diff --git a/.github/workflows/pipeline.yml b/.github/workflows/pipeline.yml index aa8e29e49..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 @@ -549,7 +570,7 @@ jobs: env: GH_TOKEN: ${{ github.token }} run: | - for artifact in $(gh api repos/${{ github.repository }}/actions/artifacts | jq -r '.artifacts[] | select(.name | startswith("digests-")) | .id'); do + for artifact in $(gh api repos/${{ github.repository }}/actions/runs/${{ github.run_id }}/artifacts | jq -r '.artifacts[] | select(.name | startswith("digests-")) | .id'); do gh api --method DELETE repos/${{ github.repository }}/actions/artifacts/$artifact done 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/Makefile b/Makefile index 81a609422..f31d98a01 100644 --- a/Makefile +++ b/Makefile @@ -20,7 +20,11 @@ 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 +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/adapters/gotaglib/end_to_end_test.go b/adapters/gotaglib/end_to_end_test.go index e7dd18ac1..0f9a90d94 100644 --- a/adapters/gotaglib/end_to_end_test.go +++ b/adapters/gotaglib/end_to_end_test.go @@ -90,7 +90,7 @@ var _ = Describe("Extractor", func() { info.FileInfo = testFileInfo{FileInfo: fileInfo} metadata := metadata.New(path, info) - return new(metadata.ToMediaFile(1, "folderID")) + return new(metadata.ToMediaFile(model.Library{ID: 1}, "folderID")) } BeforeEach(func() { diff --git a/adapters/lastfm/agent.go b/adapters/lastfm/agent.go index 7f005db1a..dbd73c30d 100644 --- a/adapters/lastfm/agent.go +++ b/adapters/lastfm/agent.go @@ -241,6 +241,10 @@ func (l *lastfmAgent) GetSimilarSongsByTrack(ctx context.Context, id, name, arti var ( artistOpenGraphQuery = cascadia.MustCompile(`html > head > meta[property="og:image"]`) artistIgnoredImage = "2a96cbd8b46e442fc41c2b86b821562f" // Last.fm artist placeholder image name + + // Not a RetryLaterError on purpose: parking the agent would also stall its API-backed + // methods, which the page block does not affect. + errNoArtistPage = errors.New("no artist image in Last.fm page") ) func (l *lastfmAgent) GetArtistImages(ctx context.Context, _, name, mbid string) ([]agents.ExternalImage, error) { @@ -267,7 +271,9 @@ func (l *lastfmAgent) GetArtistImages(ctx context.Context, _, name, mbid string) var res []agents.ExternalImage n := cascadia.Query(node, artistOpenGraphQuery) if n == nil { - return res, nil + // A real artist page always has og:image; its absence means a bot challenge or a redesign. + log.Warn(ctx, "Last.fm did not return a usable artist page", "name", name, "url", a.URL) + return nil, errNoArtistPage } for _, attr := range n.Attr { if attr.Key != "content" { diff --git a/adapters/lastfm/agent_test.go b/adapters/lastfm/agent_test.go index ce81e0916..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) @@ -648,18 +648,41 @@ var _ = Describe("lastfmAgent", func() { Expect(images).To(BeEmpty()) }) - It("returns empty list if page has no meta tags", func() { + It("errors when the page has no meta tags", func() { fApi, _ := os.Open("tests/fixtures/lastfm.artist.getinfo.json") apiClient.Res = http.Response{Body: fApi, StatusCode: 200} fScraper, _ := os.Open("tests/fixtures/lastfm.artist.page.no_meta.html") httpClient.Res = http.Response{Body: fScraper, StatusCode: 200} + _, err := agent.GetArtistImages(ctx, "123", "U2", "") + Expect(err).To(MatchError(errNoArtistPage)) + }) + + It("errors when Last.fm serves a bot challenge page", func() { + fApi, _ := os.Open("tests/fixtures/lastfm.artist.getinfo.json") + apiClient.Res = http.Response{Body: fApi, StatusCode: 200} + + fScraper, _ := os.Open("tests/fixtures/lastfm.artist.page.challenge.html") + httpClient.Res = http.Response{Body: fScraper, StatusCode: 200} + images, err := agent.GetArtistImages(ctx, "123", "U2", "") - Expect(err).ToNot(HaveOccurred()) + Expect(err).To(MatchError(errNoArtistPage)) Expect(images).To(BeEmpty()) }) + It("does not park the agent: the failure is not a retry-later", func() { + // A RetryLaterError would cool down the agent's API-backed methods too. + fApi, _ := os.Open("tests/fixtures/lastfm.artist.getinfo.json") + apiClient.Res = http.Response{Body: fApi, StatusCode: 200} + + fScraper, _ := os.Open("tests/fixtures/lastfm.artist.page.challenge.html") + httpClient.Res = http.Response{Body: fScraper, StatusCode: 200} + + _, err := agent.GetArtistImages(ctx, "123", "U2", "") + Expect(errors.Is(err, agents.ErrRetryLater)).To(BeFalse()) + }) + It("returns error if API call fails", func() { apiClient.Err = errors.New("api error") _, err := agent.GetArtistImages(ctx, "123", "U2", "") diff --git a/adapters/lastfm/auth_router.go b/adapters/lastfm/auth_router.go index 411bf069a..ec1bb9290 100644 --- a/adapters/lastfm/auth_router.go +++ b/adapters/lastfm/auth_router.go @@ -132,7 +132,7 @@ func (s *Router) callback(w http.ResponseWriter, r *http.Request) { func (s *Router) fetchSessionKey(ctx context.Context, uid, token string) error { sessionKey, err := s.client.getSession(ctx, token) if err != nil { - log.Error(ctx, "Could not fetch LastFM session key", "userId", uid, "token", token, + log.Error(ctx, "Could not fetch LastFM session key", "userId", uid, "requestId", middleware.GetReqID(ctx), err) return err } 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/lastfm/client.go b/adapters/lastfm/client.go index e468aa638..7b2601cf9 100644 --- a/adapters/lastfm/client.go +++ b/adapters/lastfm/client.go @@ -49,6 +49,12 @@ type client struct { hc httpDoer } +// escapePlus works around Last.fm decoding artist.* and track.* params twice, turning "+" into a space. +// album.getInfo decodes only once, so it must not use this. +func escapePlus(s string) string { + return strings.ReplaceAll(s, "+", "%2B") +} + func (c *client) albumGetInfo(ctx context.Context, name string, artist string, mbid string, lang string) (*Album, error) { params := url.Values{} params.Add("method", "album.getInfo") @@ -66,7 +72,7 @@ func (c *client) albumGetInfo(ctx context.Context, name string, artist string, m func (c *client) artistGetInfo(ctx context.Context, name string, lang string) (*Artist, error) { params := url.Values{} params.Add("method", "artist.getInfo") - params.Add("artist", name) + params.Add("artist", escapePlus(name)) params.Add("lang", lang) response, err := c.makeRequest(ctx, http.MethodGet, params, false) if err != nil { @@ -78,7 +84,7 @@ func (c *client) artistGetInfo(ctx context.Context, name string, lang string) (* func (c *client) artistGetSimilar(ctx context.Context, name string, limit int) (*SimilarArtists, error) { params := url.Values{} params.Add("method", "artist.getSimilar") - params.Add("artist", name) + params.Add("artist", escapePlus(name)) params.Add("limit", strconv.Itoa(limit)) response, err := c.makeRequest(ctx, http.MethodGet, params, false) if err != nil { @@ -90,7 +96,7 @@ func (c *client) artistGetSimilar(ctx context.Context, name string, limit int) ( func (c *client) artistGetTopTracks(ctx context.Context, name string, limit int) (*TopTracks, error) { params := url.Values{} params.Add("method", "artist.getTopTracks") - params.Add("artist", name) + params.Add("artist", escapePlus(name)) params.Add("limit", strconv.Itoa(limit)) response, err := c.makeRequest(ctx, http.MethodGet, params, false) if err != nil { @@ -102,8 +108,8 @@ func (c *client) artistGetTopTracks(ctx context.Context, name string, limit int) func (c *client) trackGetSimilar(ctx context.Context, name, artist string, limit int) (*SimilarTracks, error) { params := url.Values{} params.Add("method", "track.getSimilar") - params.Add("track", name) - params.Add("artist", artist) + params.Add("track", escapePlus(name)) + params.Add("artist", escapePlus(artist)) params.Add("limit", strconv.Itoa(limit)) response, err := c.makeRequest(ctx, http.MethodGet, params, false) if err != nil { diff --git a/adapters/lastfm/client_test.go b/adapters/lastfm/client_test.go index 271ae1419..f72c84b91 100644 --- a/adapters/lastfm/client_test.go +++ b/adapters/lastfm/client_test.go @@ -35,6 +35,15 @@ var _ = Describe("client", func() { Expect(album.Name).To(Equal("Believe")) Expect(httpClient.SavedRequest.URL.String()).To(Equal(apiBaseUrl + "?album=Believe&api_key=API_KEY&artist=U2&format=json&lang=pt&mbid=mbid-1234&method=album.getInfo")) }) + + It("does not double-encode plus signs", func() { + f, _ := os.Open("tests/fixtures/lastfm.album.getinfo.json") + httpClient.Res = http.Response{Body: f, StatusCode: 200} + + _, err := client.albumGetInfo(context.Background(), "Lungs", "Florence + the Machine", "", "en") + Expect(err).ToNot(HaveOccurred()) + Expect(httpClient.SavedRequest.URL.Query().Get("artist")).To(Equal("Florence + the Machine")) + }) }) Describe("artistGetInfo", func() { @@ -48,6 +57,15 @@ var _ = Describe("client", func() { Expect(httpClient.SavedRequest.URL.String()).To(Equal(apiBaseUrl + "?api_key=API_KEY&artist=U2&format=json&lang=pt&method=artist.getInfo")) }) + It("double-encodes plus signs in the artist name", func() { + f, _ := os.Open("tests/fixtures/lastfm.artist.getinfo.json") + httpClient.Res = http.Response{Body: f, StatusCode: 200} + + _, err := client.artistGetInfo(context.Background(), "Florence + the Machine", "en") + Expect(err).ToNot(HaveOccurred()) + Expect(httpClient.SavedRequest.URL.Query().Get("artist")).To(Equal("Florence %2B the Machine")) + }) + It("fails if Last.fm returns an http status != 200", func() { httpClient.Res = http.Response{ Body: io.NopCloser(bytes.NewBufferString(`Internal Server Error`)), @@ -107,6 +125,15 @@ var _ = Describe("client", func() { Expect(len(similar.Artists)).To(Equal(2)) Expect(httpClient.SavedRequest.URL.String()).To(Equal(apiBaseUrl + "?api_key=API_KEY&artist=U2&format=json&limit=2&method=artist.getSimilar")) }) + + It("double-encodes plus signs in the artist name", func() { + f, _ := os.Open("tests/fixtures/lastfm.artist.getsimilar.json") + httpClient.Res = http.Response{Body: f, StatusCode: 200} + + _, err := client.artistGetSimilar(context.Background(), "+44", 2) + Expect(err).ToNot(HaveOccurred()) + Expect(httpClient.SavedRequest.URL.Query().Get("artist")).To(Equal("%2B44")) + }) }) Describe("artistGetTopTracks", func() { @@ -119,6 +146,15 @@ var _ = Describe("client", func() { Expect(len(top.Track)).To(Equal(2)) Expect(httpClient.SavedRequest.URL.String()).To(Equal(apiBaseUrl + "?api_key=API_KEY&artist=U2&format=json&limit=2&method=artist.getTopTracks")) }) + + It("double-encodes plus signs in the artist name", func() { + f, _ := os.Open("tests/fixtures/lastfm.artist.gettoptracks.json") + httpClient.Res = http.Response{Body: f, StatusCode: 200} + + _, err := client.artistGetTopTracks(context.Background(), "C+C Music Factory", 2) + Expect(err).ToNot(HaveOccurred()) + Expect(httpClient.SavedRequest.URL.Query().Get("artist")).To(Equal("C%2BC Music Factory")) + }) }) Describe("trackGetSimilar", func() { @@ -135,6 +171,17 @@ var _ = Describe("client", func() { Expect(httpClient.SavedRequest.URL.String()).To(Equal(apiBaseUrl + "?api_key=API_KEY&artist=Depeche+Mode&format=json&limit=5&method=track.getSimilar&track=Just+Can%27t+Get+Enough")) }) + It("double-encodes plus signs in the track and artist names", func() { + f, _ := os.Open("tests/fixtures/lastfm.track.getsimilar.json") + httpClient.Res = http.Response{Body: f, StatusCode: 200} + + _, err := client.trackGetSimilar(context.Background(), "1+1", "Queen + Paul Rodgers", 5) + Expect(err).ToNot(HaveOccurred()) + query := httpClient.SavedRequest.URL.Query() + Expect(query.Get("track")).To(Equal("1%2B1")) + Expect(query.Get("artist")).To(Equal("Queen %2B Paul Rodgers")) + }) + It("returns empty list when no similar tracks found", func() { f, _ := os.Open("tests/fixtures/lastfm.track.getsimilar.unknown.json") httpClient.Res = http.Response{Body: f, StatusCode: 200} 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/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/scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go b/api/api_suite_test.go similarity index 69% rename from scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go rename to api/api_suite_test.go index 815940381..62a547b7c 100644 --- a/scanner/metadata_old/ffmpeg/ffmpeg_suite_test.go +++ b/api/api_suite_test.go @@ -1,4 +1,4 @@ -package ffmpeg +package api_test import ( "testing" @@ -9,9 +9,9 @@ import ( . "github.com/onsi/gomega" ) -func TestFFMpeg(t *testing.T) { - tests.Init(t, true) +func TestAPI(t *testing.T) { + tests.Init(t, false) log.SetLevel(log.LevelFatal) RegisterFailHandler(Fail) - RunSpecs(t, "FFMpeg Suite") + 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/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/backup.go b/cmd/backup.go index c02f3a19f..932a833f2 100644 --- a/cmd/backup.go +++ b/cmd/backup.go @@ -2,9 +2,7 @@ package cmd import ( "context" - "fmt" - "os" - "strings" + "path/filepath" "time" "github.com/navidrome/navidrome/conf" @@ -31,7 +29,7 @@ func init() { pruneCmd.Flags().BoolVarP(&force, "force", "f", false, "bypass warning when backup count is zero") backupRoot.AddCommand(pruneCmd) - restoreCommand.Flags().StringVarP(&restorePath, "backup-file", "b", "", "path of backup database to restore") + restoreCommand.Flags().StringVarP(&restorePath, "backup-file", "b", "", "file name of the backup database to restore (resolved against the backup directory unless it is an absolute path)") restoreCommand.Flags().BoolVarP(&force, "force", "f", false, "bypass restore warning") _ = restoreCommand.MarkFlagRequired("backup-file") backupRoot.AddCommand(restoreCommand) @@ -78,24 +76,12 @@ func runBackup(ctx context.Context) { conf.Server.Backup.Path = conf.NewDir(backupDir) } - idx := strings.LastIndex(conf.Server.DbPath, "?") - var path string - - if idx == -1 { - path = conf.Server.DbPath - } else { - path = conf.Server.DbPath[:idx] - } - - if _, err := os.Stat(path); os.IsNotExist(err) { - log.Fatal("No existing database", "path", path) - return - } + requireExistingDB() start := time.Now() path, err := db.Backup(ctx) if err != nil { - log.Fatal("Error backing up database", "backup path", conf.Server.BasePath, err) + log.Fatal("Error backing up database", "backupPath", conf.Server.Backup.Path, err) } elapsed := time.Since(start) @@ -111,36 +97,17 @@ func runPrune(ctx context.Context) { conf.Server.Backup.Count = backupCount } - if conf.Server.Backup.Count == 0 && !force { - fmt.Println("Warning: pruning ALL backups") - fmt.Printf("Please enter YES (all caps) to continue: ") - var input string - _, err := fmt.Scanln(&input) - - if input != "YES" || err != nil { - log.Warn("Prune cancelled") - return - } - } - - idx := strings.LastIndex(conf.Server.DbPath, "?") - var path string - - if idx == -1 { - path = conf.Server.DbPath - } else { - path = conf.Server.DbPath[:idx] - } - - if _, err := os.Stat(path); os.IsNotExist(err) { - log.Fatal("No existing database", "path", path) + if conf.Server.Backup.Count == 0 && !force && !confirmYES("Warning: pruning ALL backups") { + log.Warn("Prune cancelled") return } + requireExistingDB() + start := time.Now() count, err := db.Prune(ctx) if err != nil { - log.Fatal("Error pruning up database", "backup path", conf.Server.BasePath, err) + log.Fatal("Error pruning database", "backupPath", conf.Server.Backup.Path, err) } elapsed := time.Since(start) @@ -149,36 +116,29 @@ func runPrune(ctx context.Context) { } func runRestore(ctx context.Context) { - idx := strings.LastIndex(conf.Server.DbPath, "?") - var path string + requireExistingDB() - if idx == -1 { - path = conf.Server.DbPath - } else { - path = conf.Server.DbPath[:idx] - } - - if _, err := os.Stat(path); os.IsNotExist(err) { - log.Fatal("No existing database", "path", path) - return - } - - if !force { - fmt.Println("Warning: restoring the Navidrome database should only be done offline, especially if your backup is very old.") - fmt.Printf("Please enter YES (all caps) to continue: ") - var input string - _, err := fmt.Scanln(&input) - - if input != "YES" || err != nil { - log.Warn("Restore cancelled") + // A relative --backup-file is resolved against Backup.Path, the same folder + // `backup create` writes to. Without this, the value was treated as relative + // to the working directory, where the file does not exist. + if !filepath.IsAbs(restorePath) { + backupPath, err := conf.Server.Backup.Path.Path() + if err != nil { + log.Fatal("Backup directory not available", "backupPath", conf.Server.Backup.Path, err) return } + restorePath = filepath.Join(backupPath, restorePath) + } + + if !force && !confirmYES("Warning: restoring the Navidrome database should only be done offline, especially if your backup is very old.") { + log.Warn("Restore cancelled") + return } start := time.Now() err := db.Restore(ctx, restorePath) if err != nil { - log.Fatal("Error restoring database", "backup path", conf.Server.BasePath, err) + log.Fatal("Error restoring database", "backupFile", restorePath, err) } elapsed := time.Since(start) diff --git a/cmd/doctor.go b/cmd/doctor.go new file mode 100644 index 000000000..fd0eb1cd7 --- /dev/null +++ b/cmd/doctor.go @@ -0,0 +1,99 @@ +package cmd + +import ( + "context" + "database/sql" + "fmt" + "io" + "os" + + "github.com/navidrome/navidrome/db" + "github.com/spf13/cobra" +) + +func init() { + rootCmd.AddCommand(doctorCmd) +} + +var doctorCmd = &cobra.Command{ + Use: "doctor", + Short: "Check your Navidrome installation for problems", + Long: "Run read-only health checks and report what was found. Checks the database for " + + "corruption and foreign key violations, and reports whether 'navidrome search rebuild' " + + "can fix what it finds. This command never alters your data", + Run: func(cmd *cobra.Command, _ []string) { + runDoctor(cmd.Context()) + }, +} + +func runDoctor(ctx context.Context) { + requireExistingDB() + + healthy := doctor(ctx, db.Db(), os.Stdout) + db.Close(ctx) + if !healthy { + os.Exit(1) + } +} + +const recoveryAdvice = "Restore a backup (navidrome backup restore), or try SQLite's '.recover' command." + +func printFindings(out io.Writer, check, noun string, items []string) { + fmt.Fprintf(out, "%s reported %d %s:\n", check, len(items), noun) + for _, item := range items { + fmt.Fprintln(out, " "+item) + } +} + +func doctor(ctx context.Context, database *sql.DB, out io.Writer) bool { + healthy := true + + fmt.Fprintln(out, "Checking database integrity...") + issues, truncated, err := db.IntegrityCheck(ctx, database) + switch { + case err != nil: + fmt.Fprintln(out, "The integrity check could not complete: "+err.Error()) + fmt.Fprintln(out, recoveryAdvice) + return false + case len(issues) == 0: + fmt.Fprintln(out, "Integrity check passed.") + default: + healthy = false + printFindings(out, "Integrity check", "issue(s)", issues) + switch { + case truncated: + fmt.Fprintln(out, "The integrity check stopped at its limit, so the damage may reach further than listed.") + fmt.Fprintln(out, recoveryAdvice) + case db.IsFTSCorruptionOnly(issues): + fmt.Fprintln(out, "Corruption is limited to the search index. Run 'navidrome search rebuild' to fix it.") + default: + fmt.Fprintln(out, "Corruption is not limited to the search index, and cannot be repaired automatically.") + fmt.Fprintln(out, recoveryAdvice) + } + } + + fmt.Fprintln(out, "Checking foreign keys...") + violations, err := db.ForeignKeyCheck(ctx, database) + switch { + case err != nil: + healthy = false + fmt.Fprintln(out, "The foreign key check could not complete: "+err.Error()) + case len(violations) == 0: + fmt.Fprintln(out, "Foreign key check passed.") + default: + healthy = false + lines := make([]string, 0, len(violations)) + for _, v := range violations { + lines = append(lines, + fmt.Sprintf("%s: %d row(s) reference missing rows in %s", v.Table, v.Count, v.Parent)) + } + printFindings(out, "Foreign key check", "violation(s)", lines) + fmt.Fprintln(out, "These are orphaned rows, not corruption. 'navidrome scan -f' clears some of them "+ + "in library data; the rest have to be removed by hand.") + } + + if healthy { + fmt.Fprintln(out, "Database is healthy.") + } + return healthy +} diff --git a/cmd/doctor_test.go b/cmd/doctor_test.go new file mode 100644 index 000000000..7f8cd9728 --- /dev/null +++ b/cmd/doctor_test.go @@ -0,0 +1,124 @@ +package cmd + +import ( + "context" + "database/sql" + "os" + "path/filepath" + "strings" + + "github.com/navidrome/navidrome/db" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("doctor", func() { + var ( + ctx context.Context + dbPath string + database *sql.DB + out *strings.Builder + reopen func() + ) + + // A file-backed DB so specs can corrupt raw pages; a table named like a real FTS + // search table so IsFTSCorruptionOnly matches, plus a parent/child pair for FK checks. + BeforeEach(func() { + ctx = context.Background() + dbPath = filepath.Join(GinkgoT().TempDir(), "doctor.db") + reopen = func() { + var err error + database, err = sql.Open(db.Dialect, dbPath) + Expect(err).ToNot(HaveOccurred()) + database.SetMaxOpenConns(1) + } + reopen() + DeferCleanup(func() { _ = database.Close() }) + + for _, stmt := range []string{ + `create virtual table media_file_fts using fts5(title, content='', content_rowid='rowid')`, + `insert into media_file_fts(rowid, title) values (1, 'teenage lobotomy'), (2, 'rockaway beach')`, + `create table library(id integer primary key)`, + `create table media_file(id integer primary key, library_id integer references library(id))`, + } { + _, err := database.ExecContext(ctx, stmt) + Expect(err).ToNot(HaveOccurred()) + } + out = &strings.Builder{} + }) + + It("reports a healthy database", func() { + Expect(doctor(ctx, database, out)).To(BeTrue()) + Expect(out.String()).To(ContainSubstring("Database is healthy.")) + }) + + It("points to 'search rebuild' when corruption is limited to the search index", func() { + _, err := database.ExecContext(ctx, + `update media_file_fts_data set block = x'deadbeefdeadbeef' where id > 1`) + Expect(err).ToNot(HaveOccurred()) + + Expect(doctor(ctx, database, out)).To(BeFalse()) + Expect(out.String()).To(ContainSubstring("navidrome search rebuild")) + }) + + It("points to a backup restore when corruption is not limited to the search index", func() { + _, err := database.ExecContext(ctx, + `insert into library(id) + with recursive s(x) as (select 1 union all select x+1 from s where x < 200) + select x from s`) + Expect(err).ToNot(HaveOccurred()) + var rootPage, pageSize int64 + Expect(database.QueryRowContext(ctx, + `select rootpage from sqlite_master where name = 'library'`).Scan(&rootPage)).To(Succeed()) + Expect(database.QueryRowContext(ctx, `pragma page_size`).Scan(&pageSize)).To(Succeed()) + Expect(database.Close()).To(Succeed()) + f, err := os.OpenFile(dbPath, os.O_WRONLY, 0600) + Expect(err).ToNot(HaveOccurred()) + _, err = f.WriteAt([]byte{0xde, 0xad, 0xbe, 0xef, 0xde, 0xad, 0xbe, 0xef}, (rootPage-1)*pageSize+40) + Expect(err).ToNot(HaveOccurred()) + Expect(f.Close()).To(Succeed()) + reopen() + + Expect(doctor(ctx, database, out)).To(BeFalse()) + Expect(out.String()).To(ContainSubstring("backup restore")) + Expect(out.String()).ToNot(ContainSubstring("search rebuild")) + }) + + It("reports foreign key violations", func() { + _, err := database.ExecContext(ctx, `pragma foreign_keys = off`) + Expect(err).ToNot(HaveOccurred()) + _, err = database.ExecContext(ctx, `insert into media_file(id, library_id) values (1, 999)`) + Expect(err).ToNot(HaveOccurred()) + + Expect(doctor(ctx, database, out)).To(BeFalse()) + Expect(out.String()).To(ContainSubstring("Foreign key check reported")) + Expect(out.String()).To(ContainSubstring("media_file")) + Expect(out.String()).To(ContainSubstring("navidrome scan -f")) + // GC never touches player, share or playqueue, so don't promise a full cleanup. + Expect(out.String()).To(ContainSubstring("removed by hand")) + }) + + // Every issue names an FTS-like index, so IsFTSCorruptionOnly alone would send the + // user to 'search rebuild', but the pragma stopped at its limit without saying so. + It("does not blame the search index when the issue list is truncated", func() { + for _, stmt := range []string{ + `create table t(a, b)`, + `with recursive s(x) as (select 1 union all select x+1 from s where x < 300) + insert into t select x, x + 10000 from s`, + `create index media_file_fts_probe on t(a)`, + `pragma writable_schema=on`, + `update sqlite_master set sql = 'CREATE INDEX media_file_fts_probe ON t(b)' + where name = 'media_file_fts_probe'`, + } { + _, err := database.ExecContext(ctx, stmt) + Expect(err).ToNot(HaveOccurred()) + } + Expect(database.Close()).To(Succeed()) + reopen() + + Expect(doctor(ctx, database, out)).To(BeFalse()) + Expect(out.String()).ToNot(ContainSubstring("search rebuild")) + Expect(out.String()).To(ContainSubstring("backup restore")) + }) +}) diff --git a/cmd/inspect.go b/cmd/inspect.go index 5e88793cc..05f569f3e 100644 --- a/cmd/inspect.go +++ b/cmd/inspect.go @@ -1,13 +1,17 @@ package cmd import ( + "context" "encoding/json" "fmt" + "path/filepath" "strings" "github.com/navidrome/navidrome/core" + "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/persistence" "github.com/pelletier/go-toml/v2" "github.com/spf13/cobra" "gopkg.in/yaml.v3" @@ -28,7 +32,7 @@ var inspectCmd = &cobra.Command{ Long: "Show file tags as seen by Navidrome", Args: cobra.MinimumNArgs(1), Run: func(cmd *cobra.Command, args []string) { - runInspector(args) + runInspector(cmd.Context(), args) }, } @@ -55,18 +59,24 @@ func prettyMarshal(v any) ([]byte, error) { return []byte(res.String()), nil } -func runInspector(args []string) { +func runInspector(ctx context.Context, args []string) { marshal := marshalers[format] if marshal == nil { log.Fatal("Invalid format", "format", format) } + libs := loadLibraries(ctx) + matcher := model.NewLibraryMatcher(libs) var out []core.InspectOutput for _, filePath := range args { if !model.IsAudioFile(filePath) { log.Warn("Not an audio file", "file", filePath) continue } - output, err := core.Inspect(filePath, 1, "") + lib, ok := libraryForFile(matcher, filePath) + if !ok && len(libs) > 0 { + log.Warn("File is not in any library, using the global PID config", "file", filePath) + } + output, err := core.Inspect(filePath, lib, "") if err != nil { log.Warn("Unable to process file", "file", filePath, "error", err) continue @@ -77,3 +87,33 @@ func runInspector(args []string) { data, _ := marshal(out) fmt.Println(string(data)) } + +// loadLibraries reads the libraries, so each file gets its library's PID config. It never creates a DB. +func loadLibraries(ctx context.Context) model.Libraries { + if dbFile, ok := existingDBFile(); !ok { + log.Warn(ctx, "No database found, using the global PID config", "path", dbFile) + return nil + } + defer db.Init(ctx)() + libs, err := persistence.New(db.Db()).Library().GetAll(ctx) + if err != nil { + log.Warn(ctx, "Could not load libraries, using the global PID config", err) + return nil + } + for i := range libs { + if absPath, err := filepath.Abs(libs[i].Path); err == nil { + libs[i].Path = absPath + } + } + return libs +} + +// libraryForFile falls back to the default library with no overrides, which uses the global PID config. +func libraryForFile(matcher *model.LibraryMatcher, filePath string) (model.Library, bool) { + if absPath, err := filepath.Abs(filePath); err == nil { + if lib, ok := matcher.FindLibrary(absPath); ok { + return lib, true + } + } + return model.Library{ID: model.DefaultLibraryID}, false +} diff --git a/cmd/inspect_test.go b/cmd/inspect_test.go new file mode 100644 index 000000000..728dc8770 --- /dev/null +++ b/cmd/inspect_test.go @@ -0,0 +1,61 @@ +package cmd + +import ( + "os" + "path/filepath" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("inspect", func() { + Describe("libraryForFile", func() { + var matcher *model.LibraryMatcher + var root string + + BeforeEach(func() { + root = GinkgoT().TempDir() + cwd, err := os.Getwd() + Expect(err).ToNot(HaveOccurred()) + matcher = model.NewLibraryMatcher(model.Libraries{ + {ID: 1, Path: filepath.Join(root, "music")}, + {ID: 2, Path: filepath.Join(cwd, "loose"), PIDAlbum: "folder"}, + }) + }) + + It("returns the library that contains an absolute path", func() { + lib, ok := libraryForFile(matcher, filepath.Join(root, "music", "album", "track.mp3")) + Expect(ok).To(BeTrue()) + Expect(lib.ID).To(Equal(1)) + }) + + It("resolves a relative path against the working directory", func() { + lib, ok := libraryForFile(matcher, filepath.Join("loose", "track.mp3")) + Expect(ok).To(BeTrue()) + Expect(lib.PIDAlbum).To(Equal("folder")) + }) + + It("falls back to the default library without overrides", func() { + lib, ok := libraryForFile(matcher, filepath.Join(root, "elsewhere", "track.mp3")) + Expect(ok).To(BeFalse()) + Expect(lib).To(Equal(model.Library{ID: model.DefaultLibraryID})) + }) + }) + + Describe("loadLibraries", func() { + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + }) + + It("does not create a database when there is none", func() { + dbFile := filepath.Join(GinkgoT().TempDir(), "navidrome.db") + conf.Server.DbPath = dbFile + "?_journal_mode=WAL" + + Expect(loadLibraries(GinkgoT().Context())).To(BeNil()) + Expect(dbFile).ToNot(BeAnExistingFile()) + }) + }) +}) diff --git a/cmd/missing.go b/cmd/missing.go new file mode 100644 index 000000000..ce95e39f0 --- /dev/null +++ b/cmd/missing.go @@ -0,0 +1,171 @@ +package cmd + +import ( + "bufio" + "context" + "encoding/csv" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "strconv" + "strings" + + "github.com/Masterminds/squirrel" + "github.com/navidrome/navidrome/core" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/utils/slice" + "github.com/spf13/cobra" +) + +var missingListFormat string + +func init() { + missingListCmd.Flags().StringVarP(&missingListFormat, "format", "f", "csv", "output format [supported values: csv, json]") + missingCmd.AddCommand(missingListCmd) + missingCmd.AddCommand(missingFixCmd) + rootCmd.AddCommand(missingCmd) +} + +var ( + missingCmd = &cobra.Command{ + Use: "missing", + Short: "Manage missing files", + Long: "List files marked as missing and remap them onto existing files", + } + + missingListCmd = &cobra.Command{ + Use: "list", + Short: "List missing files", + Run: func(cmd *cobra.Command, _ []string) { + runMissingList(cmd.Context()) + }, + } + + missingFixCmd = &cobra.Command{ + Use: "fix ", + Short: "Remap a missing file onto an existing file", + Long: "Remap a file marked as missing onto an existing (non-missing) file, the same way\n" + + "the scanner reconciles moved or renamed files. Each argument may be a media file ID,\n" + + "a library-relative path, or a libraryID:path pair.", + Args: cobra.ExactArgs(2), + Run: func(cmd *cobra.Command, args []string) { + runMissingFix(cmd.Context(), args[0], args[1]) + }, + } +) + +type displayMissingFile struct { + ID string `json:"id"` + LibraryID int `json:"libraryId"` + Title string `json:"title"` + Album string `json:"album"` + Artist string `json:"artist"` + Path string `json:"path"` +} + +func runMissingList(ctx context.Context) { + if missingListFormat != "csv" && missingListFormat != "json" { + log.Fatal("Invalid output format. Must be one of csv, json", "format", missingListFormat) + } + + ds, ctx := getAdminContext(ctx) + mfs, err := ds.MediaFile().GetCursor(ctx, model.QueryOptions{ + Filters: squirrel.Eq{"missing": true}, + Sort: "path", + }) + if err == nil { + err = writeMissingList(os.Stdout, missingListFormat, mfs) + } + if err != nil { + log.Fatal(ctx, "Failed to retrieve missing files", err) + } +} + +// writeMissingList streams the cursor so a library with many missing files doesn't get loaded into memory +func writeMissingList(w io.Writer, format string, mfs model.MediaFileCursor) error { + if format == "json" { + bw := bufio.NewWriter(w) + _, _ = io.WriteString(bw, "[") + sep := "" + for mf, err := range mfs { + if err != nil { + return err + } + j, _ := json.Marshal(displayMissingFile{ID: mf.ID, LibraryID: mf.LibraryID, Title: mf.Title, Album: mf.Album, Artist: mf.Artist, Path: mf.Path}) + _, _ = fmt.Fprintf(bw, "%s%s", sep, j) + sep = "," + } + _, _ = io.WriteString(bw, "]\n") + return bw.Flush() + } + + cw := csv.NewWriter(w) + _ = cw.Write([]string{"id", "library id", "title", "album", "artist", "path"}) + for mf, err := range mfs { + if err != nil { + return err + } + _ = cw.Write([]string{mf.ID, strconv.Itoa(mf.LibraryID), mf.Title, mf.Album, mf.Artist, mf.Path}) + } + cw.Flush() + return cw.Error() +} + +func runMissingFix(ctx context.Context, missingRef, targetRef string) { + ds, ctx := getAdminContext(ctx) + + missing := resolveMediaFile(ctx, ds, missingRef) + target := resolveMediaFile(ctx, ds, targetRef) + + if err := core.NewMaintenance(ds).RemapMissingFile(ctx, missing.ID, target.ID); err != nil { + log.Fatal(ctx, "Failed to remap missing file", "missing", missing.Path, "target", target.Path, err) + } + fmt.Printf("Remapped %q onto %q\n", missing.Path, target.Path) +} + +// 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().Get(ctx, ref) + if err == nil { + return mf + } + if !errors.Is(err, model.ErrNotFound) { + log.Fatal(ctx, "Error looking up media file", "ref", ref, err) + } + + mfs, err := ds.MediaFile().FindByPaths(ctx, []string{ref}) + if err != nil { + log.Fatal(ctx, "Error looking up media file by path", "ref", ref, err) + } + if len(mfs) == 0 { + log.Fatal(ctx, "No media file found", "ref", ref) + } + mfs = preferQualified(ref, mfs) + if len(mfs) > 1 { + log.Fatal(ctx, "Path matches multiple files; disambiguate with an ID or libraryID:path", "ref", ref, "matches", len(mfs)) + } + return &mfs[0] +} + +// preferQualified resolves the ambiguity FindByPaths creates by searching a "libraryID:path" +// reference both ways: an explicit library wins over a file literally named like one. +func preferQualified(ref string, mfs model.MediaFiles) model.MediaFiles { + id, path, ok := strings.Cut(ref, ":") + if !ok { + return mfs + } + libraryID, err := strconv.Atoi(id) + if err != nil { + return mfs + } + qualified := slice.Filter(mfs, func(mf model.MediaFile) bool { + return mf.LibraryID == libraryID && strings.EqualFold(mf.Path, path) + }) + if len(qualified) == 0 { + return mfs + } + return qualified +} diff --git a/cmd/missing_test.go b/cmd/missing_test.go new file mode 100644 index 000000000..2e96cf3c7 --- /dev/null +++ b/cmd/missing_test.go @@ -0,0 +1,81 @@ +package cmd + +import ( + "errors" + "strings" + + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("writeMissingList", func() { + cursor := func(err error, mfs ...model.MediaFile) model.MediaFileCursor { + return func(yield func(model.MediaFile, error) bool) { + for _, mf := range mfs { + if !yield(mf, nil) { + return + } + } + if err != nil { + yield(model.MediaFile{}, err) + } + } + } + song := model.MediaFile{ID: "1", LibraryID: 1, Path: "Bach: Goldberg/01.mp3", Title: "Aria", Album: "Goldberg", Artist: "Bach"} + + It("writes csv with a header, quoting as needed", func() { + var out strings.Builder + Expect(writeMissingList(&out, "csv", cursor(nil, song))).To(Succeed()) + Expect(out.String()).To(Equal("id,library id,title,album,artist,path\n1,1,Aria,Goldberg,Bach,Bach: Goldberg/01.mp3\n")) + }) + + It("writes a json array", func() { + var out strings.Builder + Expect(writeMissingList(&out, "json", cursor(nil, song, song))).To(Succeed()) + Expect(out.String()).To(MatchJSON(`[ + {"id":"1","libraryId":1,"path":"Bach: Goldberg/01.mp3","title":"Aria","album":"Goldberg","artist":"Bach"}, + {"id":"1","libraryId":1,"path":"Bach: Goldberg/01.mp3","title":"Aria","album":"Goldberg","artist":"Bach"} + ]`)) + }) + + It("writes an empty json array when nothing is missing", func() { + var out strings.Builder + Expect(writeMissingList(&out, "json", cursor(nil))).To(Succeed()) + Expect(out.String()).To(MatchJSON(`[]`)) + }) + + It("returns the cursor's error", func() { + var out strings.Builder + Expect(writeMissingList(&out, "csv", cursor(errors.New("boom"), song))).To(MatchError("boom")) + }) +}) + +var _ = Describe("preferQualified", func() { + target := model.MediaFile{ID: "want", LibraryID: 1, Path: "foo.mp3"} + decoy := model.MediaFile{ID: "decoy", LibraryID: 1, Path: "1:foo.mp3"} + + It("picks the library-qualified match over a literal path that looks like one", func() { + Expect(preferQualified("1:foo.mp3", model.MediaFiles{target, decoy})).To(Equal(model.MediaFiles{target})) + }) + + It("picks the named library when the same path exists in two", func() { + other := model.MediaFile{ID: "other", LibraryID: 2, Path: "foo.mp3"} + Expect(preferQualified("1:foo.mp3", model.MediaFiles{target, other})).To(Equal(model.MediaFiles{target})) + }) + + It("leaves an unqualified reference ambiguous", func() { + both := model.MediaFiles{target, {ID: "other", LibraryID: 2, Path: "foo.mp3"}} + Expect(preferQualified("foo.mp3", both)).To(Equal(both)) + }) + + It("leaves it alone when the prefix is not a library id", func() { + both := model.MediaFiles{decoy, {ID: "other", LibraryID: 2, Path: "1:foo.mp3"}} + Expect(preferQualified("x:foo.mp3", both)).To(Equal(both)) + }) + + It("leaves it alone when no candidate matches the qualified form", func() { + both := model.MediaFiles{decoy, {ID: "other", LibraryID: 2, Path: "1:foo.mp3"}} + Expect(preferQualified("9:nope.mp3", both)).To(Equal(both)) + }) +}) 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 c4e360010..e39c55365 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,16 +78,17 @@ 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)) g.Go(startPlaybackServer(ctx)) + g.Go(startJellyfinDiscovery(ctx)) g.Go(schedulePeriodicBackup(ctx)) g.Go(startInsightsCollector(ctx)) g.Go(scheduleDBAnalyzer(ctx)) @@ -101,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. @@ -132,6 +137,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 @@ -182,46 +190,50 @@ func schedulePeriodicScan(ctx context.Context) func() error { } } -func pidHashChanged(ds model.DataStore) (bool, error) { - pidAlbum, err := ds.Property(context.Background()).DefaultGet(consts.PIDAlbumKey, "") +// librariesWithChangedPID returns the names of the libraries whose effective PID config differs from +// the one used by their last finished scan +func librariesWithChangedPID(ctx context.Context, ds model.DataStore) ([]string, error) { + libs, err := ds.Library().GetAll(ctx) if err != nil { - return false, err + return nil, err } - pidTrack, err := ds.Property(context.Background()).DefaultGet(consts.PIDTrackKey, "") - if err != nil { - return false, err + var names []string + for _, lib := range libs { + if lib.PIDChanged() { + names = append(names, lib.Name) + } } - return !strings.EqualFold(pidAlbum, conf.Server.PID.Album) || !strings.EqualFold(pidTrack, conf.Server.PID.Track), nil + return names, nil } // runInitialScan runs an initial scan of the music library if needed. 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 } - pidHasChanged, err := pidHashChanged(ds) + pidChangedLibs, err := librariesWithChangedPID(ctx, ds) if err != nil { return err } scanOnStartup := conf.Server.Scanner.Enabled && conf.Server.Scanner.ScanOnStartup - scanNeeded := scanOnStartup || inProgress || fullScanRequired == "1" || pidHasChanged + scanNeeded := scanOnStartup || inProgress || fullScanRequired == "1" || len(pidChangedLibs) > 0 time.Sleep(2 * time.Second) // Wait 2 seconds before the initial scan if scanNeeded { s := CreateScanner(ctx) switch { case fullScanRequired == "1": log.Warn(ctx, "Full scan required after migration") - _ = ds.Property(ctx).Delete(consts.FullScanAfterMigrationFlagKey) - case pidHasChanged: - log.Warn(ctx, "PID config changed, performing full scan") - fullScanRequired = "1" + _ = ds.Property().Delete(ctx, consts.FullScanAfterMigrationFlagKey) + case len(pidChangedLibs) > 0: + // Includes never-scanned libraries. The scanner rescans in full only the ones that need it + log.Warn(ctx, "Libraries with a new or changed PID config, scanning", "libraries", pidChangedLibs) case inProgress: log.Warn(ctx, "Resuming interrupted scan") default: @@ -343,6 +355,18 @@ func startInsightsCollector(ctx context.Context) func() error { } } +// startJellyfinDiscovery never returns an error: a discovery failure must not stop the server. +func startJellyfinDiscovery(ctx context.Context) func() error { + return func() error { + if !conf.Server.Jellyfin.Enabled || !conf.Server.Jellyfin.AutoDiscovery { + log.Debug("Jellyfin auto-discovery is DISABLED") + return nil + } + CreateJellyfinDiscovery().Serve(ctx) + return nil + } +} + // startPlaybackServer starts the Navidrome playback server, if configured. // It is responsible for the Jukebox functionality func startPlaybackServer(ctx context.Context) func() error { @@ -362,10 +386,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 { @@ -373,26 +413,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/cmd/root_test.go b/cmd/root_test.go index af8d44e7e..423cd2a8f 100644 --- a/cmd/root_test.go +++ b/cmd/root_test.go @@ -1,6 +1,7 @@ package cmd import ( + "errors" "net/http" "net/http/httptest" "path" @@ -9,6 +10,8 @@ import ( "github.com/go-chi/chi/v5" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -44,3 +47,30 @@ var _ = Describe("profilerHandler", func() { Entry("with a trailing-slash BasePath", "/music/"), ) }) + +var _ = Describe("librariesWithChangedPID", func() { + var ds *tests.MockDataStore + var libs *tests.MockLibraryRepo + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + libs = &tests.MockLibraryRepo{} + ds = &tests.MockDataStore{MockedLibrary: libs} + }) + + It("returns only the libraries whose PID config changed", func() { + pid := model.Library{}.EffectivePID() + libs.SetData(model.Libraries{ + {ID: 1, Name: "Same", ScannedPIDAlbum: pid.Album, ScannedPIDTrack: pid.Track}, + {ID: 2, Name: "Changed", PIDAlbum: "folder", ScannedPIDAlbum: pid.Album, ScannedPIDTrack: pid.Track}, + {ID: 3, Name: "Never scanned"}, + }) + Expect(librariesWithChangedPID(GinkgoT().Context(), ds)).To(ConsistOf("Changed", "Never scanned")) + }) + + It("returns the error from the repository", func() { + libs.Err = errors.New("db down") + _, err := librariesWithChangedPID(GinkgoT().Context(), ds) + Expect(err).To(MatchError("db down")) + }) +}) diff --git a/cmd/search.go b/cmd/search.go new file mode 100644 index 000000000..46b7eebd5 --- /dev/null +++ b/cmd/search.go @@ -0,0 +1,55 @@ +package cmd + +import ( + "context" + "fmt" + + "github.com/navidrome/navidrome/db" + "github.com/navidrome/navidrome/log" + "github.com/spf13/cobra" +) + +var searchRebuildForce bool + +func init() { + rootCmd.AddCommand(searchRoot) + + searchRebuildCmd.Flags().BoolVarP(&searchRebuildForce, "force", "f", false, "bypass rebuild confirmation") + searchRoot.AddCommand(searchRebuildCmd) +} + +var ( + searchRoot = &cobra.Command{ + Use: "search", + Short: "Search index maintenance", + } + + searchRebuildCmd = &cobra.Command{ + Use: "rebuild", + Short: "Rebuild the full-text search index", + Long: "Drop and rebuild the full-text search index from the library data. Fixes a corrupted " + + "or desynced search index without any data loss. Note that 'navidrome doctor' detects a " + + "corrupted index, but cannot tell when the index has merely drifted out of sync with the " + + "library. This must be done offline", + Run: func(cmd *cobra.Command, _ []string) { + runSearchRebuild(cmd.Context()) + }, + } +) + +func runSearchRebuild(ctx context.Context) { + requireExistingDB() + + if !searchRebuildForce && !confirmYES("This will rebuild the search index. Make sure Navidrome is not running.") { + log.Warn("Rebuild cancelled") + return + } + + fmt.Println("Rebuilding the search index...") + err := db.RebuildFTS(ctx, db.Db()) + db.Close(ctx) + if err != nil { + log.Fatal("Error rebuilding the search index", err) + } + fmt.Println("Search index rebuilt successfully.") +} diff --git a/cmd/svc.go b/cmd/svc.go index 7fec708ff..4e8b1fd85 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{} } @@ -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 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 74da51828..f35a31fb1 100644 --- a/cmd/utils.go +++ b/cmd/utils.go @@ -5,8 +5,11 @@ import ( "errors" "fmt" "io" + "os" + "strings" "text/tabwriter" + "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/core/auth" "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/log" @@ -15,6 +18,28 @@ import ( "github.com/navidrome/navidrome/persistence" ) +// existingDBFile returns the database file (DbPath minus DSN params), and whether it exists. +func existingDBFile() (string, bool) { + path, _, _ := strings.Cut(conf.Server.DbPath, "?") + _, err := os.Stat(path) + return path, err == nil +} + +// requireExistingDB aborts the command when the database file does not exist. +func requireExistingDB() { + if path, ok := existingDBFile(); !ok { + log.Fatal("No existing database", "path", path) + } +} + +func confirmYES(warning string) bool { + fmt.Println(warning) + fmt.Printf("Please enter YES (all caps) to continue: ") + var input string + _, err := fmt.Scanln(&input) + return input == "YES" && err == nil +} + // newTabWriter keeps every CLI table on the same column settings. func newTabWriter(out io.Writer) *tabwriter.Writer { return tabwriter.NewWriter(out, 0, 4, 2, ' ', 0) @@ -32,14 +57,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/cmd/wire_gen.go b/cmd/wire_gen.go index de4c55a1e..2a396689d 100644 --- a/cmd/wire_gen.go +++ b/cmd/wire_gen.go @@ -21,6 +21,7 @@ import ( "github.com/navidrome/navidrome/core/metrics" "github.com/navidrome/navidrome/core/playback" "github.com/navidrome/navidrome/core/playlists" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/core/scrobbler" "github.com/navidrome/navidrome/core/sonic" "github.com/navidrome/navidrome/core/stream" @@ -30,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" @@ -70,7 +72,7 @@ func CreateNativeAPIRouter(ctx context.Context) *nativeapi.Router { insights := metrics.GetInstance(dataStore) broker := events.GetBroker() metricsMetrics := metrics.GetPrometheusInstance(dataStore) - modelScanner := scanner.New(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) + modelScanner := scanner.GetInstance(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) watcher := scanner.GetWatcher(dataStore, modelScanner) manager := plugins.GetManager(dataStore, broker, metricsMetrics) library := core.NewLibrary(dataStore, modelScanner, watcher, broker, manager) @@ -79,7 +81,8 @@ func CreateNativeAPIRouter(ctx context.Context) *nativeapi.Router { agentsAgents := agents.GetAgents(dataStore, manager) matcherMatcher := matcher.New(dataStore) provider := external.NewProvider(dataStore, agentsAgents, matcherMatcher, broker) - router := nativeapi.New(dataStore, share, playlistsPlaylists, insights, library, user, maintenance, manager, uploader, provider) + quickConnect := quickconnect.GetInstance() + router := nativeapi.New(dataStore, share, playlistsPlaylists, insights, library, user, maintenance, manager, uploader, provider, quickConnect) return router } @@ -92,8 +95,9 @@ func CreateSubsonicAPIRouter(ctx context.Context) *subsonic.Router { artworkArtwork := artwork.NewArtwork(dataStore, fileCache, imageStore, fFmpeg) transcodingCache := stream.GetTranscodingCache() mediaStreamer := stream.NewMediaStreamer(dataStore, fFmpeg, transcodingCache) + transcodeDecider := stream.NewTranscodeDecider(dataStore, fFmpeg) share := core.NewShare(dataStore) - archiver := core.NewArchiver(mediaStreamer, dataStore, share) + archiver := core.NewArchiver(mediaStreamer, transcodeDecider, dataStore, share, artworkArtwork) players := core.NewPlayers(dataStore) broker := events.GetBroker() metricsMetrics := metrics.GetPrometheusInstance(dataStore) @@ -103,11 +107,10 @@ func CreateSubsonicAPIRouter(ctx context.Context) *subsonic.Router { provider := external.NewProvider(dataStore, agentsAgents, matcherMatcher, broker) uploader := artwork.NewUploader(dataStore) playlistsPlaylists := playlists.NewPlaylists(dataStore, uploader) - modelScanner := scanner.New(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) + modelScanner := scanner.GetInstance(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) playTracker := scrobbler.GetPlayTracker(dataStore, broker, manager) playbackServer := playback.GetInstance(dataStore) lyricsLyrics := lyrics.NewLyrics(dataStore, manager) - transcodeDecider := stream.NewTranscodeDecider(dataStore, fFmpeg) sonicSonic := sonic.New(dataStore, manager, matcherMatcher) router := subsonic.New(dataStore, artworkArtwork, mediaStreamer, archiver, players, provider, modelScanner, broker, playlistsPlaylists, playTracker, share, playbackServer, metricsMetrics, lyricsLyrics, transcodeDecider, sonicSonic) return router @@ -135,7 +138,15 @@ func CreateJellyfinAPIRouter(ctx context.Context) *jellyfin.Router { provider := external.NewProvider(dataStore, agentsAgents, matcherMatcher, broker) sonicSonic := sonic.New(dataStore, manager, matcherMatcher) lyricsLyrics := lyrics.NewLyrics(dataStore, manager) - router := jellyfin.New(dataStore, artworkArtwork, mediaStreamer, transcodeDecider, players, playTracker, playlistsPlaylists, provider, sonicSonic, lyricsLyrics, broker) + quickConnect := quickconnect.GetInstance() + router := jellyfin.New(dataStore, artworkArtwork, mediaStreamer, transcodeDecider, players, playTracker, playlistsPlaylists, provider, sonicSonic, lyricsLyrics, broker, quickConnect) + return router +} + +func CreateAPIv1Router(ctx context.Context) *apiv1.Router { + sqlDB := db.Db() + dataStore := persistence.New(sqlDB) + router := apiv1.New(dataStore) return router } @@ -148,9 +159,10 @@ func CreatePublicRouter() *public.Router { artworkArtwork := artwork.NewArtwork(dataStore, fileCache, imageStore, fFmpeg) transcodingCache := stream.GetTranscodingCache() mediaStreamer := stream.NewMediaStreamer(dataStore, fFmpeg, transcodingCache) + transcodeDecider := stream.NewTranscodeDecider(dataStore, fFmpeg) share := core.NewShare(dataStore) - archiver := core.NewArchiver(mediaStreamer, dataStore, share) - router := public.New(dataStore, artworkArtwork, mediaStreamer, share, archiver) + archiver := core.NewArchiver(mediaStreamer, transcodeDecider, dataStore, share, artworkArtwork) + router := public.New(dataStore, artworkArtwork, mediaStreamer, transcodeDecider, share, archiver) return router } @@ -168,6 +180,13 @@ func CreateListenBrainzRouter() *listenbrainz.Router { return router } +func CreateJellyfinDiscovery() *jellyfin.Discovery { + sqlDB := db.Db() + dataStore := persistence.New(sqlDB) + discovery := jellyfin.NewDiscovery(dataStore) + return discovery +} + func CreateInsights() metrics.Insights { sqlDB := db.Db() dataStore := persistence.New(sqlDB) @@ -189,7 +208,7 @@ func CreateScanner(ctx context.Context) model.Scanner { uploader := artwork.NewUploader(dataStore) playlistsPlaylists := playlists.NewPlaylists(dataStore, uploader) metricsMetrics := metrics.GetPrometheusInstance(dataStore) - modelScanner := scanner.New(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) + modelScanner := scanner.GetInstance(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) return modelScanner } @@ -200,7 +219,7 @@ func CreateScanWatcher(ctx context.Context) scanner.Watcher { uploader := artwork.NewUploader(dataStore) playlistsPlaylists := playlists.NewPlaylists(dataStore, uploader) metricsMetrics := metrics.GetPrometheusInstance(dataStore) - modelScanner := scanner.New(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) + modelScanner := scanner.GetInstance(ctx, dataStore, broker, playlistsPlaylists, metricsMetrics) watcher := scanner.GetWatcher(dataStore, modelScanner) return watcher } @@ -249,7 +268,7 @@ func getPluginManager() *plugins.Manager { // wire_injectors.go: -var allProviders = wire.NewSet(core.Set, artwork.Set, server.New, subsonic.New, jellyfin.New, nativeapi.New, public.New, persistence.New, lastfm.NewRouter, listenbrainz.NewRouter, events.GetBroker, scanner.New, 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 8bc404fb1..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" @@ -36,13 +37,15 @@ var allProviders = wire.NewSet( server.New, subsonic.New, jellyfin.New, + jellyfin.NewDiscovery, + apiv1.New, nativeapi.New, public.New, persistence.New, lastfm.NewRouter, listenbrainz.NewRouter, events.GetBroker, - scanner.New, + scanner.GetInstance, scanner.GetWatcher, metrics.GetPrometheusInstance, db.Db, @@ -90,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, @@ -108,6 +117,12 @@ func CreateListenBrainzRouter() *listenbrainz.Router { )) } +func CreateJellyfinDiscovery() *jellyfin.Discovery { + panic(wire.Build( + allProviders, + )) +} + func CreateInsights() metrics.Insights { panic(wire.Build( allProviders, diff --git a/conf/configuration.go b/conf/configuration.go index ff119417a..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 { @@ -235,6 +236,8 @@ type jellyfinOptions struct { // ExposedPublicUsers is a comma-separated list of usernames to advertise on the unauthenticated // GET /Users/Public, so Jellyfin clients can show a login user-picker. Empty exposes no users. ExposedPublicUsers string + AutoDiscovery bool + QuickConnect bool // MaxConcurrentStreams bounds how many collection responses can stream at once. Each holds a DB // cursor — and its pooled connection — for the whole client-paced response, so without a bound // enough slow clients would take the entire pool and stall the scanner, scrobbles and the UI. @@ -407,7 +410,7 @@ func Load(noConfigDump bool) { if mkErr := os.MkdirAll(filepath.Dir(Server.LogFile), os.ModePerm); mkErr != nil { logFatal(fmt.Sprintf("Error creating log file directory: %s", mkErr.Error())) } - out, err = os.OpenFile(Server.LogFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) + out, err = os.OpenFile(Server.LogFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600) if err != nil { logFatal(fmt.Sprintf("Error opening log file %s: %s", Server.LogFile, err.Error())) } @@ -512,6 +515,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) @@ -1087,6 +1095,8 @@ func setViperDefaults() { viper.SetDefault("listenbrainz.trackalgorithm", consts.DefaultListenBrainzTrackAlgorithm) viper.SetDefault("jellyfin.enabled", false) viper.SetDefault("jellyfin.servername", "") + viper.SetDefault("jellyfin.autodiscovery", false) + viper.SetDefault("jellyfin.quickconnect", true) viper.SetDefault("enablescrobblehistory", true) viper.SetDefault("httpheaders.frameoptions", "DENY") viper.SetDefault("backup.path", "") @@ -1115,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/conf/configuration_test.go b/conf/configuration_test.go index 2c7f8edaa..8c4c8ab86 100644 --- a/conf/configuration_test.go +++ b/conf/configuration_test.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "runtime" "testing" "time" @@ -329,6 +330,21 @@ var _ = Describe("Configuration", func() { }).To(PanicWith(ContainSubstring("Error creating log file directory"))) }) + It("creates the log file readable only by the owner", func() { + if runtime.GOOS == "windows" { + Skip("file modes are not enforced on Windows") + } + logFile := filepath.Join(GinkgoT().TempDir(), "navidrome.log") + viper.SetDefault("datafolder", GinkgoT().TempDir()) + viper.SetDefault("logfile", logFile) + DeferCleanup(log.SetOutput, os.Stderr) + conf.Load(true) + + info, err := os.Stat(logFile) + Expect(err).ToNot(HaveOccurred()) + Expect(info.Mode().Perm()).To(Equal(os.FileMode(0600))) + }) + It("is called when BaseURL is invalid", func() { viper.SetDefault("datafolder", GinkgoT().TempDir()) viper.SetDefault("baseurl", "://invalid") @@ -386,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/consts/consts.go b/consts/consts.go index 486ea66bc..42e9ec42f 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" @@ -59,6 +60,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. @@ -154,8 +156,6 @@ const ( //DefaultAlbumPID = "album_legacy" DefaultAlbumPID = "musicbrainz_albumid|albumartistid,album,albumversion,releasedate" DefaultTrackPID = "musicbrainz_trackid|albumid,discnumber,tracknumber,title" - PIDAlbumKey = "PIDAlbum" - PIDTrackKey = "PIDTrack" ) const ( diff --git a/contrib/navidrome.service b/contrib/navidrome.service index 5e6cbedce..ef61d2c65 100644 --- a/contrib/navidrome.service +++ b/contrib/navidrome.service @@ -36,7 +36,7 @@ RestrictNamespaces=yes RestrictRealtime=yes SystemCallFilter=@system-service SystemCallFilter=~@privileged @resources -SystemCallFilter=setrlimit +SystemCallFilter=setrlimit mbind SystemCallArchitectures=native UMask=0066 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 c9436279e..60eb44858 100644 --- a/core/archiver.go +++ b/core/archiver.go @@ -2,23 +2,32 @@ package core import ( "archive/zip" + "cmp" "context" "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 @@ -26,18 +35,20 @@ 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, decider stream.TranscodeDecider, ds model.DataStore, shares Share, artwork artwork.Artwork) Archiver { + return &archiver{ds: ds, ms: ms, decider: decider, shares: shares, artwork: artwork} } type archiver struct { - ds model.DataStore - ms stream.MediaStreamer - shares Share + ds model.DataStore + ms stream.MediaStreamer + decider stream.TranscodeDecider + 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 { @@ -47,28 +58,30 @@ 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 { - mfs, err := a.ds.MediaFile(ctx).GetAll(model.QueryOptions{Filters: filters, Sort: "album"}) +// 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().GetAll(ctx, model.QueryOptions{Filters: filters, Sort: "album"}) if err != nil { log.Error(ctx, "Error loading mediafiles from artist", "id", id, err) return err } 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) - if addErr := a.addFileToZip(ctx, z, mf, format, bitrate, file); errors.Is(addErr, stream.ErrTooManyTranscodes) { + req := a.resolveRequest(ctx, &mf, format, bitrate) + file := a.albumFilename(mf, req.Format, isMultiDisc, folder) + if addErr := a.addFileToZip(ctx, z, mf, req, file); errors.Is(addErr, stream.ErrTooManyTranscodes) { // Stop iterating: continuing would just rack up more // rejections from the limiter. Close finalises whatever // tracks were already written; the rejected one is not @@ -78,7 +91,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) @@ -96,7 +112,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 +174,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 @@ -114,27 +184,31 @@ 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 { - 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 } 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)) for idx, mf := range mfs { - file := a.playlistFilename(mf, format, idx) - if addErr := a.addFileToZip(ctx, z, mf, format, bitrate, file); errors.Is(addErr, stream.ErrTooManyTranscodes) { + req := a.resolveRequest(ctx, &mf, format, bitrate) + file := a.playlistFilename(mf, req.Format, idx) + if addErr := a.addFileToZip(ctx, z, mf, req, file); errors.Is(addErr, stream.ErrTooManyTranscodes) { // Abort the whole archive: continuing would silently emit // empty zip entries since the headers are already written. _ = z.Close() @@ -143,6 +217,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 { @@ -179,7 +254,14 @@ func (a *archiver) playlistFilename(mf model.MediaFile, format string, idx int) return fmt.Sprintf("%02d - %s - %s.%s", idx+1, str.SanitizeFilename(mf.Artist), str.SanitizeFilename(mf.Title), ext) } -func (a *archiver) addFileToZip(ctx context.Context, z *zip.Writer, mf model.MediaFile, format string, bitrate int, filename string) error { +func (a *archiver) resolveRequest(ctx context.Context, mf *model.MediaFile, format string, bitrate int) stream.Request { + if format == "" || format == "raw" { + return stream.Request{Format: "raw"} + } + return a.decider.ResolveRequest(ctx, mf, format, bitrate, 0) +} + +func (a *archiver) addFileToZip(ctx context.Context, z *zip.Writer, mf model.MediaFile, req stream.Request, filename string) error { path := mf.AbsolutePath() // Open the source before writing the zip entry header so a rejection @@ -187,13 +269,13 @@ func (a *archiver) addFileToZip(ctx context.Context, z *zip.Writer, mf model.Med // archive. var r io.ReadCloser var err error - if format != "raw" && format != "" { - r, err = a.ms.NewStream(ctx, &mf, stream.Request{Format: format, BitRate: bitrate}) + if req.Format != "raw" { + r, err = a.ms.NewStream(ctx, &mf, req) } else { r, err = os.Open(path) } if err != nil { - log.Error(ctx, "Error opening file for zipping", "file", path, "format", format, err) + log.Error(ctx, "Error opening file for zipping", "file", path, "format", req.Format, err) return err } defer func() { @@ -220,3 +302,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 461af1800..178d1b6b9 100644 --- a/core/archiver_test.go +++ b/core/archiver_test.go @@ -4,13 +4,18 @@ import ( "archive/zip" "bytes" "context" + "errors" "io" "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/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" @@ -21,15 +26,19 @@ var _ = Describe("Archiver", func() { var ( arch core.Archiver ms *mockMediaStreamer + dc *fakeDecider ds *mockDataStore sh *mockShare + ca *mockCoverArt ) BeforeEach(func() { ms = &mockMediaStreamer{} + dc = &fakeDecider{} sh = &mockShare{} ds = &mockDataStore{} - arch = core.NewArchiver(ms, ds, sh) + ca = &mockCoverArt{images: map[string][]byte{}} + arch = core.NewArchiver(ms, dc, ds, sh, ca) }) Context("ZipAlbum", func() { @@ -45,7 +54,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) @@ -59,6 +68,23 @@ var _ = Describe("Archiver", func() { Expect(zr.File[0].Name).To(Equal("Album_Promo/01 - track1.mp3")) Expect(zr.File[1].Name).To(Equal("Album_Promo/02 - track2.mp3")) }) + + It("streams the request resolved by the transcode decider and names the entry after its format", func() { + mfRepo := &mockMediaFileRepository{} + mfRepo.On("GetAll", mock.Anything).Return(model.MediaFiles{{Path: "test_data/01 - track1.flac", Suffix: "flac", AlbumID: "1"}}, nil) + ds.On("MediaFile").Return(mfRepo) + resolved := stream.Request{Format: "opus", BitRate: 128, SampleRate: 48000, Channels: 2} + dc.resolved = &resolved + ms.On("NewStream", mock.Anything, mock.Anything, resolved).Return(io.NopCloser(strings.NewReader("test")), nil).Once() + + out := new(bytes.Buffer) + Expect(arch.ZipAlbum(GinkgoT().Context(), "1", "mp3", 128, out)).To(Succeed()) + ms.AssertExpectations(GinkgoT()) + + zr, err := zip.NewReader(bytes.NewReader(out.Bytes()), int64(out.Len())) + Expect(err).ToNot(HaveOccurred()) + Expect(zr.File[0].Name).To(HaveSuffix("01 - track1.opus")) + }) }) Context("ZipArtist", func() { @@ -77,7 +103,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) @@ -91,6 +117,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() { @@ -105,7 +265,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() @@ -155,6 +315,30 @@ var _ = Describe("Archiver", func() { }) Context("ZipPlaylist", func() { + It("names the entries and the M3U lines after the resolved format", func() { + pls := &model.Playlist{ID: "1", Name: "Test Playlist", Tracks: []model.PlaylistTrack{ + {MediaFile: model.MediaFile{Path: "test_data/01 - track1.flac", Suffix: "flac", Artist: "Artist 1", Title: "track1"}}, + }} + plRepo := &mockPlaylistRepository{} + plRepo.On("GetWithTracks", "1", true, false).Return(pls, nil) + ds.On("Playlist").Return(plRepo) + dc.resolved = &stream.Request{Format: "opus", BitRate: 128} + ms.On("NewStream", mock.Anything, mock.Anything, *dc.resolved).Return(io.NopCloser(strings.NewReader("test")), nil) + + out := new(bytes.Buffer) + Expect(arch.ZipPlaylist(GinkgoT().Context(), "1", "mp3", 128, out)).To(Succeed()) + + zr, err := zip.NewReader(bytes.NewReader(out.Bytes()), int64(out.Len())) + Expect(err).ToNot(HaveOccurred()) + Expect(zr.File[0].Name).To(Equal("01 - Artist 1 - track1.opus")) + m3u, err := zr.File[1].Open() + Expect(err).ToNot(HaveOccurred()) + defer m3u.Close() + content, err := io.ReadAll(m3u) + Expect(err).ToNot(HaveOccurred()) + Expect(string(content)).To(ContainSubstring("01 - Artist 1 - track1.opus")) + }) + It("zips a playlist correctly", func() { tracks := []model.PlaylistTrack{ {MediaFile: model.MediaFile{Path: "test_data/01 - track1.mp3", Suffix: "mp3", AlbumID: "1", Album: "Album 1", DiscNumber: 1, Artist: "AC/DC", Title: "track1"}}, @@ -169,7 +353,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) @@ -196,24 +380,195 @@ 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 } -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{} } @@ -222,7 +577,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 } @@ -231,7 +586,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) } @@ -241,7 +596,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) } @@ -259,6 +614,19 @@ func (m *mockMediaStreamer) NewStream(ctx context.Context, mf *model.MediaFile, return &stream.Stream{ReadCloser: args.Get(0).(io.ReadCloser)}, nil } +// fakeDecider echoes the legacy format/bitrate unless a resolved request is set. +type fakeDecider struct { + stream.TranscodeDecider + resolved *stream.Request +} + +func (f *fakeDecider) ResolveRequest(_ context.Context, _ *model.MediaFile, format string, bitRate int, offset int) stream.Request { + if f.resolved != nil { + return *f.resolved + } + return stream.Request{Format: format, BitRate: bitRate, Offset: offset} +} + type mockShare struct { mock.Mock core.Share diff --git a/core/artwork/artwork.go b/core/artwork/artwork.go index 663d06d25..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) @@ -168,7 +168,12 @@ func (s *service) serveHash(ctx context.Context, artID model.ArtworkID, ia *mode if !entityExists(ctx, s.ds, artID) { return nil, ErrUnavailable } - art, err := s.ds.Artwork(ctx).GetImage(ia.Hash) + // Checked here, not in openOriginal: a resize-cache hit never opens the source. + if isFileBacked(ia.Source) && !model.IsImageFile(ia.SourcePath) { + 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().GetImage(ctx, ia.Hash) if err != nil { if errors.Is(err, model.ErrNotFound) { return s.dangling(ctx, artID) @@ -259,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) @@ -278,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 } @@ -337,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 1ea82b7fa..921a77b89 100644 --- a/core/artwork/artwork_suite_test.go +++ b/core/artwork/artwork_suite_test.go @@ -1,19 +1,23 @@ package artwork import ( + "context" "io/fs" + "net/netip" "net/url" "os" "path/filepath" "runtime" "strings" "testing" + "time" "github.com/navidrome/navidrome/core/storage" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/metadata" "github.com/navidrome/navidrome/tests" + "github.com/navidrome/navidrome/utils/httpclient" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" "go.uber.org/goleak" @@ -37,6 +41,14 @@ func TestArtwork(t *testing.T) { RunSpecs(t, "Artwork Suite") } +// productionImageClient keeps the guarded client for the specs that assert it refuses loopback. +var productionImageClient = remoteImageClient + +// httptest servers listen on loopback, which the production client refuses. +var _ = BeforeSuite(func() { + remoteImageClient = httpclient.NewExternal(5*time.Second, netip.MustParsePrefix("127.0.0.0/8"), netip.MustParsePrefix("::1/128")) +}) + // osDirFS wraps os.DirFS as a storage.MusicFS for integration tests. type osDirFS struct{ fs.FS } @@ -97,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 907b300de..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()) @@ -159,13 +159,60 @@ var _ = Describe("Artwork", func() { Expect(readAll(img)).To(Equal(coverBytes)) }) + It("treats a file-backed row pointing at a non-image file as dangling", func() { + dir := GinkgoT().TempDir() + secretPath := filepath.Join(dir, "config.ini") + Expect(os.WriteFile(secretPath, []byte("password=secret"), 0600)).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "dddddddddddddddd", Mime: "image/jpeg"})).To(Succeed()) + seedEntity("al", "alni") + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ + ItemKind: "al", ItemID: "alni", Hash: "dddddddddddddddd", + Source: "folder", SourcePath: secretPath, RefMtime: fileMtime(secretPath), + })).To(Succeed()) + + _, err := svc.Get(ctx, model.MustParseArtworkID("al-alni"), 0, false) + Expect(err).To(MatchError(ErrUnavailable)) + Expect(queueRepo.Data[primaryKey("al", "alni")].Priority).To(Equal(model.ArtworkPriorityScan)) + }) + + It("refuses a non-image file-backed row even when a resized copy is already cached", func() { + secret := []byte("password=secret") + dir := GinkgoT().TempDir() + secretPath := filepath.Join(dir, "config.ini") + Expect(os.WriteFile(secretPath, secret, 0600)).To(Succeed()) + Expect(artRepo.PutImage(ctx, &model.Artwork{Hash: "eeeeeeeeeeeeeeee", Mime: "image/jpeg"})).To(Succeed()) + seedEntity("al", "alnic") + Expect(artRepo.PutItemArtwork(ctx, &model.ItemArtwork{ + ItemKind: "al", ItemID: "alnic", Hash: "eeeeeeeeeeeeeeee", + Source: "folder", SourcePath: secretPath, RefMtime: fileMtime(secretPath), + })).To(Succeed()) + + // Older versions cached the raw bytes when the resize failed. + seed := func() (io.ReadCloser, error) { return io.NopCloser(bytes.NewReader(secret)), nil } + stream, err := imgCache.Get(ctx, &resizedItem{hash: "eeeeeeeeeeeeeeee", size: 100, open: seed, ffmpeg: ffm}) + Expect(err).ToNot(HaveOccurred()) + Expect(io.ReadAll(stream)).To(Equal(secret)) + Expect(stream.Close()).To(Succeed()) + Eventually(func(g Gomega) { + s, err := imgCache.Get(ctx, &resizedItem{hash: "eeeeeeeeeeeeeeee", size: 100, ffmpeg: ffm, + open: func() (io.ReadCloser, error) { return nil, os.ErrNotExist }}) + g.Expect(err).ToNot(HaveOccurred()) + g.Expect(s.Cached).To(BeTrue()) + _ = s.Close() + }).Should(Succeed()) + + _, err = svc.Get(ctx, model.MustParseArtworkID("al-alnic"), 100, false) + Expect(err).To(MatchError(ErrUnavailable)) + Expect(queueRepo.Data[primaryKey("al", "alnic")].Priority).To(Equal(model.ArtworkPriorityScan)) + }) + It("treats a full-size mtime mismatch as dangling: unavailable, re-enqueued at Scan, state untouched", 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()) @@ -173,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")) }) @@ -182,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()) @@ -205,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()) @@ -225,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)) }) @@ -236,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)) }) }) @@ -263,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) @@ -296,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)) }) @@ -447,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()) @@ -488,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 5a9396cc3..38c28a9b2 100644 --- a/core/artwork/e2e/artist_test.go +++ b/core/artwork/e2e/artist_test.go @@ -90,6 +90,28 @@ var _ = Describe("Artist artwork resolution", func() { }) }) + When("the artist's only album folder has no images of its own", func() { + // Artist/ + // ├── backdrop1.jpg + // ├── folder.jpg ← matched by folder.* + // ├── logo.png + // └── Album/ + // ├── 01 - Track.mp3 + // └── 02 - Track.mp3 + It("resolves the artist folder, not the library root", func() { + conf.Server.ArtistArtPriority = "folder.*, artist.*, album/artist.*" + setLayout(fstest.MapFS{ + "Artist/Album/01 - Track.mp3": trackFile(1, "Track 1", map[string]any{"albumartist": "Artist", "album": "Album"}), + "Artist/Album/02 - Track.mp3": trackFile(2, "Track 2", map[string]any{"albumartist": "Artist", "album": "Album"}), + "Artist/backdrop1.jpg": smallPNG("backdrop"), + "Artist/folder.jpg": smallPNG("artist-folder"), + "Artist/logo.png": smallPNG("logo"), + }) + scan() + expectArtistFolder(soleArtist(), "Artist/folder.jpg") + }) + }) + When("the artist's only album has its tracks in disc subfolders", func() { // Artist/ // ├── artist.jpg ← wins (artist.* before album/artist.*) @@ -179,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")) @@ -257,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 8881da3a6..d03e22228 100644 --- a/core/artwork/e2e/e2e_suite_test.go +++ b/core/artwork/e2e/e2e_suite_test.go @@ -56,7 +56,8 @@ func runWorkerUntil(ctx context.Context, worker *artwork.Worker, until func() bo runCtx, cancel := context.WithCancel(ctx) done := make(chan error, 1) go func() { done <- worker.Run(runCtx) }() - Eventually(until, 5*time.Second, 10*time.Millisecond).Should(BeTrue()) + // Long enough for one retry (3-7s backoff, 5s poll tick). + Eventually(until, 15*time.Second, 10*time.Millisecond).Should(BeTrue()) cancel() Eventually(done, 2*time.Second).Should(Receive(BeNil())) } @@ -66,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 1ca1ce034..3403935db 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,14 +168,15 @@ 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().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/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 } -} 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/image_store.go b/core/artwork/image_store.go index 5dbe727e4..2ded3e636 100644 --- a/core/artwork/image_store.go +++ b/core/artwork/image_store.go @@ -98,6 +98,9 @@ func (s *ImageStore) Write(hash, mimeType string, r io.Reader) error { if err := tmp.Close(); err != nil { return err } + if err := os.Chmod(tmp.Name(), 0640); err != nil { + return err + } return os.Rename(tmp.Name(), dst) } diff --git a/core/artwork/image_store_test.go b/core/artwork/image_store_test.go index f7b2467ca..7de3ba761 100644 --- a/core/artwork/image_store_test.go +++ b/core/artwork/image_store_test.go @@ -48,6 +48,17 @@ var _ = Describe("ImageStore", func() { Expect(got).To(Equal(data)) }) + It("writes group-readable image files", func() { + tests.SkipOnWindows("uses Unix file permission bits") + data := []byte("jpeg-bytes") + h, _ := hashImage(bytes.NewReader(data)) + Expect(store.Write(h, "image/jpeg", bytes.NewReader(data))).To(Succeed()) + + info, err := os.Stat(store.path(h, "image/jpeg")) + Expect(err).ToNot(HaveOccurred()) + Expect(info.Mode().Perm()).To(Equal(os.FileMode(0640))) + }) + It("is idempotent on duplicate writes and preserves the original content", func() { data := []byte("dup") h, _ := hashImage(bytes.NewReader(data)) 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/resize.go b/core/artwork/resize.go index 9536c3bb1..9b3ff46e5 100644 --- a/core/artwork/resize.go +++ b/core/artwork/resize.go @@ -76,7 +76,7 @@ func toFastScaleType(img image.Image) image.Image { } func resizeStaticImage(data []byte, size int, square bool) (io.Reader, int, error) { - original, format, err := image.Decode(bytes.NewReader(data)) + original, format, err := decodeCapped(data) if err != nil { return nil, 0, err } diff --git a/core/artwork/resize_test.go b/core/artwork/resize_test.go new file mode 100644 index 000000000..76b95974d --- /dev/null +++ b/core/artwork/resize_test.go @@ -0,0 +1,13 @@ +package artwork + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("resizeStaticImage", func() { + It("rejects images whose declared dimensions exceed the pixel cap before decoding", func() { + _, _, err := resizeStaticImage(pngHeaderWithDims(9000, 9000), 300, false) + Expect(err).To(MatchError(ContainSubstring("exceed pixel cap"))) + }) +}) diff --git a/core/artwork/resolve.go b/core/artwork/resolve.go index e7a2d3765..fb07332fe 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,11 @@ 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()"}) + tracks := r.ds.Playlist().Tracks(ctx, pl.ID, false) + if tracks == nil { + return resolution{}, fmt.Errorf("resolvePlaylist: could not load tracks for playlist %s", pl.ID) + } + albumIDs, err := tracks.GetAlbumIDs(ctx, model.QueryOptions{Max: PlaylistGridSamples, Sort: "random()"}) if err != nil { return resolution{}, err } @@ -428,7 +431,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 +442,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 } @@ -572,7 +575,7 @@ func resolveArtistFolderPattern(ctx context.Context, lib libraryView, artistFold // resolveLocalFile opens an absolute path directly. A missing path is "no source"; any other // open failure says nothing about whether the image exists. func resolveLocalFile(path, source string) (resolution, bool) { - if path == "" { + if path == "" || !model.IsImageFile(path) { return resolution{}, false } f, err := os.Open(path) diff --git a/core/artwork/resolve_test.go b/core/artwork/resolve_test.go index 402a11363..2a36531bb 100644 --- a/core/artwork/resolve_test.go +++ b/core/artwork/resolve_test.go @@ -498,6 +498,23 @@ var _ = Describe("resolveItem", func() { Expect(res.refMtime).To(BeNumerically(">", 0)) }) + It("never opens a local ExternalImageURL that is not an image file", func() { + folderRepo.result = nil // no grid tiles, so only the local file could produce a reader + dir := GinkgoT().TempDir() + secretPath := filepath.Join(dir, "config.ini") + Expect(os.WriteFile(secretPath, []byte("password=secret"), 0600)).To(Succeed()) + + plRepo := tests.CreateMockPlaylistRepo() + plRepo.SetData(model.Playlists{{ID: "plni", Name: "Playlist", ExternalImageURL: secretPath}}) + plRepo.TracksRepo = &tests.MockPlaylistTrackRepo{AlbumIDs: []string{"t1"}} + ds.MockedPlaylist = plRepo + + res, err := newResolver(ds, ag, ffm, nil).resolve(ctx, model.ArtworkQueueItem{ItemKind: "pl", ItemID: "plni"}) + Expect(err).ToNot(HaveOccurred()) + Expect(res.reader).To(BeNil()) + Expect(res.sourcePath).ToNot(Equal(secretPath)) + }) + It("routes ExternalImageURL through extGate and sets extError on transient failure", func() { conf.Server.EnableM3UExternalAlbumArt = true folderRepo.result = nil // no grid tiles, so the external failure is what surfaces @@ -690,6 +707,16 @@ var _ = Describe("resolveItem", func() { Expect(err).To(HaveOccurred()) Expect(res).To(Equal(resolution{})) }) + + It("returns an error when the playlist tracks cannot be loaded", func() { + plRepo := tests.CreateMockPlaylistRepo() + plRepo.SetData(model.Playlists{{ID: "pl4", Name: "Playlist"}}) + ds.MockedPlaylist = plRepo + + res, err := newResolver(ds, ag, ffm, nil).resolve(ctx, model.ArtworkQueueItem{ItemKind: "pl", ItemID: "pl4"}) + Expect(err).To(HaveOccurred()) + Expect(res).To(Equal(resolution{})) + }) }) }) diff --git a/core/artwork/sources.go b/core/artwork/sources.go index f2abf9da5..daa084c16 100644 --- a/core/artwork/sources.go +++ b/core/artwork/sources.go @@ -18,6 +18,7 @@ import ( "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils/httpclient" + "github.com/navidrome/navidrome/utils/netguard" "go.senan.xyz/taglib" ) @@ -37,7 +38,7 @@ func fromExternalFile(ctx context.Context, libFS fs.FS, files []string, pattern log.Warn(ctx, "Artwork: Error matching cover art file to pattern", "pattern", pattern, "file", file) continue } - if !match { + if !match || !model.IsImageFile(name) { continue } f, err := libFS.Open(file) @@ -150,10 +151,18 @@ type readCloser struct { io.Closer } +// remoteImageClient fetches URLs from playlists and agents (plugins included), so it must not reach +// internal hosts. Shared so fetches reuse connections. +var remoteImageClient = httpclient.NewExternal(5 * time.Second) + func fromURL(ctx context.Context, imageUrl *url.URL) (io.ReadCloser, string, error) { - hc := httpclient.New(5 * time.Second) req, _ := http.NewRequestWithContext(ctx, http.MethodGet, imageUrl.String(), nil) - resp, err := hc.Do(req) //nolint:gosec + resp, err := remoteImageClient.Do(req) + if errors.Is(err, netguard.ErrPrivateAddress) { + // Retrying cannot change where the URL points: settle absent instead of tripping the breaker. + log.Warn(ctx, "Artwork: Refused to fetch image from a private or loopback address", "url", imageUrl, err) + return nil, "", model.ErrNotFound + } if err != nil { return nil, "", err } diff --git a/core/artwork/sources_internal_test.go b/core/artwork/sources_internal_test.go index 4282575a5..bc81b56e7 100644 --- a/core/artwork/sources_internal_test.go +++ b/core/artwork/sources_internal_test.go @@ -5,13 +5,87 @@ import ( "errors" "io" "io/fs" + "net/http" + "net/http/httptest" + "net/netip" + "net/url" "os" + "strings" + "sync/atomic" "testing/fstest" + "time" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/utils/httpclient" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) +var _ = Describe("fromURL", func() { + var ( + hits atomic.Int32 + target *httptest.Server + ) + + BeforeEach(func() { + hits.Store(0) + target = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + hits.Add(1) + _, _ = w.Write([]byte("image-bytes")) + })) + DeferCleanup(target.Close) + }) + + useClient := func(c *http.Client) { + prev := remoteImageClient + remoteImageClient = c + DeferCleanup(func() { remoteImageClient = prev }) + } + fetch := func(rawURL string) ([]byte, error) { + u, err := url.Parse(rawURL) + Expect(err).ToNot(HaveOccurred()) + r, _, err := fromURL(GinkgoT().Context(), u) + if err != nil { + return nil, err + } + defer r.Close() + return io.ReadAll(r) + } + // Stand-in for a public host: only 127.0.0.1 is allowed, so every other private address stays refused. + onlyLocalhostV4 := func() *http.Client { + return httpclient.NewExternal(5*time.Second, netip.MustParsePrefix("127.0.0.1/32")) + } + + DescribeTable("refuses private and loopback targets as a definitive miss", + func(rawURL string) { + useClient(productionImageClient) + u, _ := url.Parse(target.URL) + _, err := fetch(strings.ReplaceAll(rawURL, "PORT", u.Port())) + Expect(err).To(MatchError(model.ErrNotFound)) + Expect(hits.Load()).To(BeZero()) + }, + Entry("IPv4 loopback", "http://127.0.0.1:PORT/x"), + Entry("localhost", "http://localhost:PORT/x"), + Entry("cloud metadata", "http://169.254.169.254/"), + Entry("IPv6 loopback", "http://[::1]/"), + ) + + It("refuses a redirect from an allowed host to a loopback address", func() { + useClient(onlyLocalhostV4()) + redirector := httptest.NewServer(http.RedirectHandler(strings.Replace(target.URL, "127.0.0.1", "127.0.0.2", 1), http.StatusFound)) + DeferCleanup(redirector.Close) + + _, err := fetch(redirector.URL) + Expect(err).To(MatchError(model.ErrNotFound)) + Expect(hits.Load()).To(BeZero()) + }) + + It("fetches from an allowed address", func() { + useClient(onlyLocalhostV4()) + Expect(fetch(target.URL + "/cover.jpg")).To(Equal([]byte("image-bytes"))) + }) +}) + var _ = Describe("fromExternalFile", func() { It("opens a matching file via the library FS", func() { fsys := fstest.MapFS{ @@ -48,6 +122,18 @@ var _ = Describe("fromExternalFile", func() { Expect(b).To(Equal([]byte("a"))) Expect(path).To(Equal("a/cover.jpg")) }) + + It("skips a matching file that is not an image", func() { + fsys := fstest.MapFS{ + "a/cover.ini": &fstest.MapFile{Data: []byte("password=secret")}, + "a/cover.jpg": &fstest.MapFile{Data: []byte("a")}, + } + f := fromExternalFile(GinkgoT().Context(), fsys, []string{"a/cover.ini", "a/cover.jpg"}, "cover.*") + r, path, err := f() + Expect(err).ToNot(HaveOccurred()) + defer r.Close() + Expect(path).To(Equal("a/cover.jpg")) + }) }) var _ = Describe("fromTag", func() { 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 bb09be55e..4be99f92e 100644 --- a/core/artwork/worker.go +++ b/core/artwork/worker.go @@ -4,9 +4,11 @@ import ( "bytes" "cmp" "context" + "fmt" "io" "math" "math/rand/v2" + "runtime/debug" "sync" "time" @@ -45,7 +47,8 @@ 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 gates map[string]*extGate @@ -59,6 +62,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 +94,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,9 +152,12 @@ 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...) + items, err := w.proc.ds.ArtworkQueue().DequeueBatch(ctx, max(16, 4*concurrency), kinds...) if err != nil { return 0, err } @@ -170,6 +182,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) @@ -231,14 +246,14 @@ func (w *Worker) process(ctx context.Context, item model.ArtworkQueueItem) (outc item.ImageType = cmp.Or(item.ImageType, model.ImageTypePrimary) trace := &ChainTrace{} ctx = withTrace(ctx, trace) - out, got, retryIn := w.proc.acquire(ctx, item) + out, got, retryIn := w.safeAcquire(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: @@ -247,7 +262,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, @@ -258,7 +273,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 @@ -266,13 +281,27 @@ 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) } } return out, got } +// safeAcquire turns a panic into a failed attempt: the drain runs on a bare goroutine, so an +// unrecovered panic would crash the server, and the still-queued row would crash it again on restart. +func (w *Worker) safeAcquire(ctx context.Context, item model.ArtworkQueueItem) (out outcome, got *acquired, retryIn time.Duration) { + defer func() { + if r := recover(); r != nil { + log.Error(ctx, "Artwork: Panic while processing item", "kind", item.ItemKind, "id", item.ItemID, + "imageType", item.ImageType, "attempts", item.Attempts, "panic", r, "stack", string(debug.Stack())) + traceStage(ctx, "panic", fmt.Errorf("%v", r)) + out, got, retryIn = outcomeFailed, nil, 0 + } + }() + return w.proc.acquire(ctx, item) +} + // recordGiveUp keeps the last failure on the state row after the queue row is deleted. An item // that never resolved has no row to update, and creating one would settle it absent. func (w *Worker) recordGiveUp(ctx context.Context, item model.ArtworkQueueItem, trace string) { @@ -280,7 +309,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) } } @@ -290,7 +319,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 ebb8de251..80ca68bc3 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,39 @@ 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) +} + +type panickingAlbumRepo struct { + *tests.MockAlbumRepo + panicID string +} + +func (r *panickingAlbumRepo) Get(ctx context.Context, id string) (*model.Album, error) { + if id == r.panicID { + panic("boom") + } + return r.MockAlbumRepo.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 +178,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 +221,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 +229,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 +244,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 +252,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 +266,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 +275,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,17 +286,47 @@ 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") }) + It("fails an item that panics, without stopping the rest of the batch", func() { + folderRepo.result = []model.Folder{{ + Path: "tests/fixtures/artist/an-album", + ImageFiles: []string{"cover.jpg"}, + }} + albums := tests.CreateMockAlbumRepo() + albums.SetData(model.Albums{ + {ID: "alboom", Name: "Album", FolderIDs: []string{"f1"}}, + {ID: "alok", Name: "Album", FolderIDs: []string{"f1"}}, + }) + ds.MockedAlbum = &panickingAlbumRepo{MockAlbumRepo: albums, panicID: "alboom"} + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alboom"})).To(Succeed()) + Expect(queueRepo.Enqueue(ctx, model.ArtworkQueueItem{ItemKind: "al", ItemID: "alok"})).To(Succeed()) + + n, err := w.drain(ctx, 1) + Expect(err).ToNot(HaveOccurred()) + Expect(n).To(Equal(2)) + + it := findQueued(queueRepo, "al", "alboom") + Expect(it).ToNot(BeNil(), "a panicking item must be rescheduled, not dropped") + Expect(it.Attempts).To(Equal(1)) + Expect(it.RetryAt).To(BeTemporally(">", time.Now())) + Expect(it.Trace).To(ContainSubstring("boom")) + + Expect(findQueued(queueRepo, "al", "alok")).To(BeNil()) + ia, err := artRepo.GetItemArtwork(ctx, model.KindAlbumArtwork, "alok", model.ImageTypePrimary) + Expect(err).ToNot(HaveOccurred()) + Expect(ia.Source).To(Equal("folder")) + }) + It("reschedules past the provider's requested delay when it exceeds the backoff", func() { conf.Server.CoverArtPriority = "external" ds.MockedAlbum.(*tests.MockAlbumRepo).SetData(model.Albums{{ID: "al9", Name: "Album"}}) // 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 +347,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 +358,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 +378,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 +388,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 +400,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 +419,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 +428,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 +436,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 +450,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 +460,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 +481,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 +496,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 +524,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 +532,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 +547,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 +574,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 +623,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 +633,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 +642,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 +776,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 +864,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 +883,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 +900,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()) } @@ -866,10 +917,38 @@ 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(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(ctx) + return n < 8 + }) + + _, err := w.drain(ctx, 1) + Expect(err).ToNot(HaveOccurred()) + + 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") + }) + 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()) } @@ -896,6 +975,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(ctx, 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/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/ffmpeg/ffmpeg.go b/core/ffmpeg/ffmpeg.go index af2dab647..af59178af 100644 --- a/core/ffmpeg/ffmpeg.go +++ b/core/ffmpeg/ffmpeg.go @@ -27,11 +27,12 @@ type TranscodeOptions struct { Command string // DB command template (used to detect custom vs default) Format string // Target format (mp3, opus, aac, flac) FilePath string - BitRate int // kbps, 0 = codec default - SampleRate int // 0 = no constraint - Channels int // 0 = no constraint - BitDepth int // 0 = no constraint; valid values: 16, 24, 32 - Offset int // seconds + BitRate int // kbps, 0 = codec default + SampleRate int // 0 = no constraint + Channels int // 0 = no constraint + BitDepth int // 0 = no constraint; valid values: 16, 24, 32 + Offset int // seconds + Duration float32 // seconds; 0 = unknown. Only used to repair a piped FLAC header. } // AudioProbeResult contains authoritative audio stream properties from ffprobe. @@ -48,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 @@ -67,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" ) @@ -86,7 +85,11 @@ func (e *ffmpeg) Transcode(ctx context.Context, opts TranscodeOptions) (io.ReadC } else { args = buildTemplateArgs(opts) } - return e.start(ctx, args) + out, err := e.start(ctx, args) + if err != nil { + return nil, err + } + return patchFLACDuration(out, opts.Duration-float32(opts.Offset)), nil } func (e *ffmpeg) ConvertAnimatedImage(ctx context.Context, reader io.Reader, maxSize int, quality int) (io.ReadCloser, error) { @@ -144,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 @@ -588,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 0fa3de111..46684fe14 100644 --- a/core/ffmpeg/ffmpeg_test.go +++ b/core/ffmpeg/ffmpeg_test.go @@ -3,6 +3,7 @@ package ffmpeg import ( "context" "errors" + "io" "os" "os/exec" "path/filepath" @@ -61,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", "-"})) }) }) @@ -684,6 +678,40 @@ var _ = Describe("ffmpeg", func() { }) Expect(err).To(MatchError(context.Canceled)) }) + + It("fills in total_samples on a piped FLAC transcode", func() { + stream, err := ff.Transcode(GinkgoT().Context(), TranscodeOptions{ + Command: "ffmpeg -i %s -map 0:a:0 -v 0 -c:a flac -f flac -", + Format: "flac", + FilePath: "tests/fixtures/test.flac", + Duration: 1, // the fixture is exactly 1s at 44100Hz + }) + Expect(err).ToNot(HaveOccurred()) + defer stream.Close() + + out, err := io.ReadAll(stream) + Expect(err).ToNot(HaveOccurred()) + Expect(string(out[:4])).To(Equal("fLaC")) + Expect(readTotalSamples(out)).To(Equal(uint64(44100))) + }) + + It("patches the duration net of the requested offset", func() { + // The command has no %t, so ffmpeg still emits the whole fixture. + // What is under test is the header arithmetic, not the audio. + stream, err := ff.Transcode(GinkgoT().Context(), TranscodeOptions{ + Command: "ffmpeg -i %s -map 0:a:0 -v 0 -c:a flac -f flac -", + Format: "flac", + FilePath: "tests/fixtures/test.flac", + Duration: 3, + Offset: 1, + }) + Expect(err).ToNot(HaveOccurred()) + defer stream.Close() + + out, err := io.ReadAll(stream) + Expect(err).ToNot(HaveOccurred()) + Expect(readTotalSamples(out)).To(Equal(uint64(2 * 44100))) + }) }) Context("stderr capture", func() { diff --git a/core/ffmpeg/flac_streaminfo.go b/core/ffmpeg/flac_streaminfo.go new file mode 100644 index 000000000..878c28718 --- /dev/null +++ b/core/ffmpeg/flac_streaminfo.go @@ -0,0 +1,66 @@ +package ffmpeg + +import ( + "bytes" + "encoding/binary" + "errors" + "io" + "math" +) + +const ( + flacPrefixLen = 26 // through the last total_samples byte + flacMaxTotalSamples = 1<<36 - 1 +) + +// patchFLACDuration fills in the STREAMINFO total_samples that ffmpeg leaves at 0 +// when writing to a pipe, since a decoder cannot seek a cached FLAC without it. +func patchFLACDuration(r io.ReadCloser, duration float32) io.ReadCloser { + if duration <= 0 { + return r + } + return &flacPatcher{ReadCloser: r, duration: duration} +} + +type flacPatcher struct { + io.ReadCloser + duration float32 + // Peeking here rather than in the constructor keeps Transcode from blocking + // until ffmpeg has emitted its first bytes. + stream io.Reader +} + +func (f *flacPatcher) Read(p []byte) (int, error) { + if f.stream == nil { + prefix := make([]byte, flacPrefixLen) + n, err := io.ReadFull(f.ReadCloser, prefix) + if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, io.ErrUnexpectedEOF) { + return 0, err + } + prefix = prefix[:n] + if err == nil { + setFLACTotalSamples(prefix, f.duration) + } + f.stream = io.MultiReader(bytes.NewReader(prefix), f.ReadCloser) + } + return f.stream.Read(p) +} + +// setFLACTotalSamples takes the rate from the header rather than the transcode +// options, so a resampled (-ar) output still gets the right count. +func setFLACTotalSamples(prefix []byte, duration float32) { + if string(prefix[:4]) != "fLaC" || prefix[4]&0x7F != 0 { + return + } + // 20-bit rate | 3-bit channels | 5-bit depth | 36-bit total_samples + info := binary.BigEndian.Uint64(prefix[18:]) + rate := info >> 44 + if rate == 0 || info&flacMaxTotalSamples != 0 { + return + } + total := math.Round(float64(duration) * float64(rate)) + if total > flacMaxTotalSamples { + return + } + binary.BigEndian.PutUint64(prefix[18:], info|uint64(total)) +} diff --git a/core/ffmpeg/flac_streaminfo_test.go b/core/ffmpeg/flac_streaminfo_test.go new file mode 100644 index 000000000..6bf3503d7 --- /dev/null +++ b/core/ffmpeg/flac_streaminfo_test.go @@ -0,0 +1,142 @@ +package ffmpeg + +import ( + "bytes" + "errors" + "io" + "os" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// Decoded independently so the specs do not mirror the production bit-twiddling. +func readSampleRate(b []byte) int { + return int(b[18])<<12 | int(b[19])<<4 | int(b[20])>>4 +} + +func readTotalSamples(b []byte) uint64 { + return uint64(b[21]&0x0F)<<32 | uint64(b[22])<<24 | uint64(b[23])<<16 | uint64(b[24])<<8 | uint64(b[25]) +} + +var _ = Describe("patchFLACDuration", func() { + var fileFLAC []byte + + // Zeroing total_samples reproduces what a piped transcode emits. + pipedFLAC := func() []byte { + b := bytes.Clone(fileFLAC) + b[21] &= 0xF0 + clear(b[22:26]) + return b + } + + readAll := func(in []byte, duration float32) []byte { + out, err := io.ReadAll(patchFLACDuration(io.NopCloser(bytes.NewReader(in)), duration)) + Expect(err).ToNot(HaveOccurred()) + return out + } + + BeforeEach(func() { + var err error + fileFLAC, err = os.ReadFile("tests/fixtures/test.flac") + Expect(err).ToNot(HaveOccurred()) + Expect(readSampleRate(fileFLAC)).To(Equal(44100)) // specs below hard-code this rate + }) + + It("fills in total_samples from the duration", func() { + out := readAll(pipedFLAC(), 1.0) + Expect(readTotalSamples(out)).To(Equal(uint64(44100))) + }) + + It("takes the sample rate from the header, not from the source file", func() { + in := pipedFLAC() + // Rewrite the header's rate to 48000, as -ar would. + in[18], in[19] = 0x0B, 0xB8 + in[20] &= 0x0F + + out := readAll(in, 2.0) + + Expect(readSampleRate(out)).To(Equal(48000)) + Expect(readTotalSamples(out)).To(Equal(uint64(96000))) + }) + + It("rounds to the nearest sample rather than truncating", func() { + // float32(0.7)*44100 is 30869.9995, so truncation would lose a sample. + out := readAll(pipedFLAC(), 0.7) + Expect(readTotalSamples(out)).To(Equal(uint64(30870))) + }) + + It("passes through when the duration overflows the 36-bit field", func() { + in := pipedFLAC() + Expect(readAll(in, 2e6)).To(Equal(in)) + }) + + It("leaves everything after the header untouched", func() { + in := pipedFLAC() + out := readAll(in, 1.0) + Expect(out).To(HaveLen(len(in))) + Expect(out[26:]).To(Equal(in[26:])) + Expect(out[:18]).To(Equal(in[:18])) + }) + + It("leaves an already-populated total_samples alone", func() { + out := readAll(fileFLAC, 99.0) + Expect(out).To(Equal(fileFLAC)) + }) + + It("passes through a stream that is not FLAC", func() { + in := []byte("ID3\x04\x00\x00\x00\x00\x00\x00 not a flac stream at all, just bytes") + Expect(readAll(in, 1.0)).To(Equal(in)) + }) + + It("passes through when the first metadata block is not STREAMINFO", func() { + in := pipedFLAC() + in[4] = 0x04 // VORBIS_COMMENT + Expect(readAll(in, 1.0)).To(Equal(in)) + }) + + It("passes through a stream shorter than the STREAMINFO fields it patches", func() { + in := pipedFLAC()[:20] + Expect(readAll(in, 1.0)).To(Equal(in)) + }) + + It("passes through an empty stream", func() { + Expect(readAll(nil, 1.0)).To(BeEmpty()) + }) + + It("passes through when the duration is zero or negative", func() { + in := pipedFLAC() + Expect(readAll(in, 0)).To(Equal(in)) + Expect(readAll(in, -5)).To(Equal(in)) + }) + + It("passes through when the header declares no sample rate", func() { + in := pipedFLAC() + in[18], in[19] = 0, 0 + in[20] &= 0x0F + Expect(readAll(in, 1.0)).To(Equal(in)) + }) + + It("propagates a read error from the underlying stream", func() { + _, err := io.ReadAll(patchFLACDuration(io.NopCloser(io.MultiReader( + bytes.NewReader(pipedFLAC()[:10]), &errReader{})), 1.0)) + Expect(err).To(MatchError("boom")) + }) + + It("closes the underlying stream", func() { + c := &closeSpy{Reader: bytes.NewReader(pipedFLAC())} + Expect(patchFLACDuration(c, 1.0).Close()).To(Succeed()) + Expect(c.closed).To(BeTrue()) + }) +}) + +type errReader struct{} + +func (e *errReader) Read([]byte) (int, error) { return 0, errors.New("boom") } + +type closeSpy struct { + io.Reader + closed bool +} + +func (c *closeSpy) Close() error { c.closed = true; return nil } diff --git a/core/inspect.go b/core/inspect.go index 01ec33760..c60459b88 100644 --- a/core/inspect.go +++ b/core/inspect.go @@ -15,7 +15,7 @@ type InspectOutput struct { MappedTags *model.MediaFile `json:"mappedTags,omitempty"` } -func Inspect(filePath string, libraryId int, folderId string) (*InspectOutput, error) { +func Inspect(filePath string, lib model.Library, folderId string) (*InspectOutput, error) { path, file := filepath.Split(filePath) s, err := storage.For(path) @@ -39,12 +39,22 @@ func Inspect(filePath string, libraryId int, folderId string) (*InspectOutput, e return nil, model.ErrNotFound } - md := metadata.New(path, tag) + md := metadata.New(scannerPath(lib, filePath), tag) result := &InspectOutput{ File: filePath, RawTags: tags[file].Tags, - MappedTags: new(md.ToMediaFile(libraryId, folderId)), + MappedTags: new(md.ToMediaFile(lib, folderId)), } return result, nil } + +// scannerPath returns the path the scanner uses for the file (relative to its library), so +// folder-based PIDs match the DB. Files outside the library keep their absolute path. +func scannerPath(lib model.Library, filePath string) string { + absPath, err := filepath.Abs(filePath) + if err != nil || lib.Path == "" { + return filePath + } + return model.LibraryRelativePath(lib.Path, absPath) +} diff --git a/core/inspect_test.go b/core/inspect_test.go new file mode 100644 index 000000000..0ac90990c --- /dev/null +++ b/core/inspect_test.go @@ -0,0 +1,45 @@ +package core_test + +import ( + "path/filepath" + + "github.com/navidrome/navidrome/core" + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Inspect", func() { + var fixtures string + + BeforeEach(func() { + var err error + fixtures, err = filepath.Abs(filepath.Join("tests", "fixtures")) + Expect(err).ToNot(HaveOccurred()) + }) + + It("maps the file with the library-relative path the scanner uses", func() { + lib := model.Library{ID: 2, Path: filepath.Dir(fixtures), PIDAlbum: "folder"} + out, err := core.Inspect(filepath.Join(fixtures, "test.mp3"), lib, "") + Expect(err).ToNot(HaveOccurred()) + Expect(out.MappedTags.Path).To(Equal("fixtures/test.mp3")) + Expect(out.MappedTags.LibraryID).To(Equal(2)) + }) + + It("gives the same IDs for relative and absolute paths", func() { + lib := model.Library{ID: 2, Path: filepath.Dir(fixtures), PIDAlbum: "folder"} + abs, err := core.Inspect(filepath.Join(fixtures, "test.mp3"), lib, "") + Expect(err).ToNot(HaveOccurred()) + rel, err := core.Inspect(filepath.Join("tests", "fixtures", "test.mp3"), lib, "") + Expect(err).ToNot(HaveOccurred()) + Expect(rel.MappedTags.AlbumID).To(Equal(abs.MappedTags.AlbumID)) + Expect(rel.MappedTags.PID).To(Equal(abs.MappedTags.PID)) + }) + + It("keeps the given path for a file outside the library", func() { + filePath := filepath.Join(fixtures, "test.mp3") + out, err := core.Inspect(filePath, model.Library{ID: model.DefaultLibraryID}, "") + Expect(err).ToNot(HaveOccurred()) + Expect(out.MappedTags.Path).To(Equal(filePath)) + }) +}) diff --git a/core/library.go b/core/library.go index d905e00cb..f1153da26 100644 --- a/core/library.go +++ b/core/library.go @@ -7,6 +7,7 @@ import ( "io/fs" "os" "path/filepath" + "slices" "strconv" "strings" "time" @@ -16,6 +17,7 @@ import ( "github.com/navidrome/navidrome/core/storage" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/metadata" "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server/events" "github.com/navidrome/navidrome/utils/slice" @@ -33,25 +35,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, + }, } } @@ -59,16 +64,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 } @@ -91,7 +96,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) } @@ -116,7 +121,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) @@ -133,25 +138,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 @@ -159,87 +153,94 @@ 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 + pidChanged := (updatesColumn(cols, "pidAlbum") && originalLib.PIDAlbum != lib.PIDAlbum) || + (updatesColumn(cols, "pidTrack") && originalLib.PIDTrack != lib.PIDTrack) - err = r.LibraryRepository.Put(lib, cols...) + err = r.LibraryRepository.Put(ctx, lib, cols...) if err != nil { return r.mapError(err) } - // 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 pathChanged && r.watcher != nil { + 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") - } + if (pathChanged || pidChanged) && r.scanner != nil { + 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{ @@ -248,7 +249,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) } @@ -256,7 +257,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) @@ -264,25 +265,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 } @@ -307,17 +308,10 @@ func (r *libraryRepositoryWrapper) mapError(err error) error { } } - switch { - case errors.Is(err, model.ErrNotFound): - return rest.ErrNotFound - case errors.Is(err, model.ErrNotAuthorized): - return rest.ErrPermissionDenied - default: - return err - } + 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 == "" { @@ -328,11 +322,20 @@ 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() } } + library.PIDAlbum = strings.TrimSpace(library.PIDAlbum) + library.PIDTrack = strings.TrimSpace(library.PIDTrack) + if err := metadata.ValidatePIDSpec(library.PIDAlbum, true); err != nil { + validationErrors["pidAlbum"] = err.Error() + } + if err := metadata.ValidatePIDSpec(library.PIDTrack, false); err != nil { + validationErrors["pidTrack"] = err.Error() + } + if len(validationErrors) > 0 { return &rest.ValidationError{Errors: validationErrors} } @@ -340,7 +343,12 @@ func (r *libraryRepositoryWrapper) validateLibrary(library *model.Library) error return nil } -func (r *libraryRepositoryWrapper) validateLibraryPath(library *model.Library) error { +// updatesColumn reports whether an update with these columns writes col. No columns means all of them. +func updatesColumn(cols []string, col string) bool { + return len(cols) == 0 || slices.Contains(cols, col) +} + +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") @@ -358,7 +366,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") } @@ -366,7 +374,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.") { @@ -393,7 +401,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 { @@ -407,13 +415,29 @@ 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) +var scanWaitInterval = time.Second + +func (r *libraryRepositoryWrapper) triggerScan(ctx context.Context, lib *model.Library, action string) { + // Runs in its own goroutine and outlives the HTTP request + ctx = context.WithoutCancel(ctx) + + // A running scan loaded the libraries before this change, and would reject a new request + for { + status, err := r.scanner.Status(ctx) + if err != nil || !status.Scanning { + break + } + time.Sleep(scanWaitInterval) + } + + 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 - if err != nil { - log.Error(r.ctx, fmt.Sprintf("Error scanning %s library", action), "libraryID", lib.ID, "name", lib.Name, err) + warnings, err := r.scanner.ScanAll(ctx, false) // Quick scan: the scanner rescans libraries with a changed PID config in full + if errors.Is(err, model.ErrAlreadyScanning) { + log.Debug(ctx, "Scan already running, it covers this change", "libraryID", lib.ID, "name", lib.Name) + } else if err != nil { + 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..e6ebb1974 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 @@ -322,12 +322,43 @@ var _ = Describe("Library Service", func() { }) }) + Describe("PID validation", func() { + pidError := func(err error, field string) string { + var validationErr *rest.ValidationError + Expect(errors.As(err, &validationErr)).To(BeTrue()) + return validationErr.Errors[field] + } + + It("rejects an unknown attribute in the album PID", func() { + _, err := repo.Save(ctx, &model.Library{Name: "Lib", Path: tempDir, PIDAlbum: "albmversion"}) + Expect(pidError(err, "pidAlbum")).To(ContainSubstring(`unknown attribute "albmversion"`)) + }) + + It("rejects albumid in the album PID", func() { + _, err := repo.Save(ctx, &model.Library{Name: "Lib", Path: tempDir, PIDAlbum: "albumid"}) + Expect(pidError(err, "pidAlbum")).To(ContainSubstring("albumid")) + }) + + It("rejects an unknown attribute in the track PID", func() { + _, err := repo.Save(ctx, &model.Library{Name: "Lib", Path: tempDir, PIDTrack: "nosuchtag"}) + Expect(pidError(err, "pidTrack")).To(ContainSubstring(`unknown attribute "nosuchtag"`)) + }) + + It("trims spaces", func() { + library := &model.Library{Name: "Lib", Path: tempDir, PIDAlbum: " folder ", PIDTrack: " "} + _, err := repo.Save(ctx, library) + Expect(err).ToNot(HaveOccurred()) + Expect(library.PIDAlbum).To(Equal("folder")) + Expect(library.PIDTrack).To(BeEmpty()) + }) + }) + Describe("Path Validation", func() { Context("Create operation", 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 +370,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 +385,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 +402,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 +424,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 +441,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 +450,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 +465,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 +477,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 +498,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 +644,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 +680,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 +701,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 @@ -679,11 +710,53 @@ var _ = Describe("Library Service", func() { }, "100ms", "10ms").Should(Equal(0)) }) + It("triggers scan when updating the library PID config", func() { + libraryRepo.SetData(model.Libraries{{ID: 1, Name: "Library", Path: tempDir}}) + + library := model.Library{ID: 1, Name: "Library", Path: tempDir, PIDAlbum: "folder"} + Expect(repo.Update(ctx, "1", library)).To(Succeed()) + + Eventually(func() int { + return scanner.GetScanAllCallCount() + }, "1s", "10ms").Should(Equal(1)) + // A quick scan: the scanner itself rescans this library in full + Expect(scanner.GetScanAllCalls()[0].FullScan).To(BeFalse()) + }) + + It("does not trigger scan when the PID fields were not sent", func() { + libraryRepo.SetData(model.Libraries{{ID: 1, Name: "Library", Path: tempDir, PIDAlbum: "folder"}}) + + // The REST layer decodes a missing pidAlbum as "". Only the sent fields count. + library := model.Library{ID: 1, Name: "Renamed", Path: tempDir} + Expect(repo.Update(ctx, "1", library, "name", "path")).To(Succeed()) + + Consistently(func() int { + return scanner.GetScanAllCallCount() + }, "100ms", "10ms").Should(Equal(0)) + }) + + It("waits for a running scan before triggering a new one", func() { + libraryRepo.SetData(model.Libraries{{ID: 1, Name: "Library", Path: tempDir}}) + scanner.SetScanning(true) + + library := model.Library{ID: 1, Name: "Library", Path: tempDir, PIDAlbum: "folder"} + Expect(repo.Update(ctx, "1", library)).To(Succeed()) + + Consistently(func() int { + return scanner.GetScanAllCallCount() + }, "200ms", "20ms").Should(Equal(0)) + + scanner.SetScanning(false) + Eventually(func() int { + return scanner.GetScanAllCallCount() + }, "3s", "20ms").Should(Equal(1)) + }) + It("does not trigger scan when library creation fails", 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 +773,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 +789,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 +804,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 +817,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 +846,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 +866,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 +881,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 +899,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 +911,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 +923,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 +936,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 +948,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 +956,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 +970,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 13d1141d3..58265ea2c 100644 --- a/core/maintenance.go +++ b/core/maintenance.go @@ -2,6 +2,7 @@ package core import ( "context" + "errors" "fmt" "slices" "sync" @@ -14,11 +15,23 @@ import ( "github.com/navidrome/navidrome/utils/slice" ) +var ( + // ErrNotMissing is returned when a remap is attempted from a file not marked as missing. + ErrNotMissing = errors.New("file is not marked as missing") + // ErrTargetMissing is returned when the remap target is itself a missing file. + ErrTargetMissing = errors.New("target file is missing") + // ErrSameFile is returned when the remap source and target are the same file. + ErrSameFile = errors.New("missing and target are the same file") +) + type Maintenance interface { // DeleteMissingFiles deletes specific missing files by their IDs DeleteMissingFiles(ctx context.Context, ids []string) error // DeleteAllMissingFiles deletes all files marked as missing DeleteAllMissingFiles(ctx context.Context) error + // RemapMissingFile moves a missing file's identity onto an existing file, the manual + // counterpart to the scanner's move detection (phaseMissingTracks.moveMatched). + RemapMissingFile(ctx context.Context, missingID, targetID string) error } type maintenanceService struct { @@ -40,6 +53,87 @@ func (s *maintenanceService) DeleteAllMissingFiles(ctx context.Context) error { return s.deleteMissing(ctx, nil) } +func (s *maintenanceService) RemapMissingFile(ctx context.Context, missingID, targetID string) error { + if missingID == targetID { + return fmt.Errorf("%w: %q", ErrSameFile, missingID) + } + + missing, err := s.ds.MediaFile().Get(ctx, missingID) + if err != nil { + return fmt.Errorf("loading missing file %q: %w", missingID, err) + } + if !missing.Missing { + return fmt.Errorf("%w: %q", ErrNotMissing, missingID) + } + + target, err := s.ds.MediaFile().GetWithParticipants(ctx, targetID) + if err != nil { + return fmt.Errorf("loading target file %q: %w", targetID, err) + } + if target.Missing { + return fmt.Errorf("%w: %q", ErrTargetMissing, targetID) + } + + oldAlbumID, newAlbumID := missing.AlbumID, target.AlbumID + + err = s.ds.WithTx(func(tx model.DataStore) error { + discardedID := target.ID + + // 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().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().ReassignReferences(ctx, discardedID, missing.ID); err != nil { + return fmt.Errorf("reassign target references: %w", err) + } + if err := tx.MediaFile().Delete(ctx, discardedID); err != nil { + return fmt.Errorf("delete discarded track: %w", err) + } + + if oldAlbumID != newAlbumID { + 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().ReassignAnnotation(ctx, oldAlbumID, newAlbumID); err != nil { + return fmt.Errorf("reassign album annotations: %w", err) + } + 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) + } + } + } + return nil + }) + if err != nil { + log.Error(ctx, "Error remapping missing file", "missing", missing.Path, "target", target.Path, err) + return err + } + + if err := s.ds.GC(ctx); err != nil { + log.Error(ctx, "Error running GC after remapping missing file", err) + return err + } + + // 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().RefreshStats(ctx, true); err != nil { + log.Error(ctx, "Error refreshing artist stats after remapping missing file", err) + } + affectedAlbumIDs := []string{newAlbumID} + if oldAlbumID != newAlbumID { + affectedAlbumIDs = append(affectedAlbumIDs, oldAlbumID) + } + if err := s.refreshAlbums(ctx, affectedAlbumIDs); err != nil { + log.Error(ctx, "Error refreshing album stats after remapping missing file", err) + } + return nil +} + // deleteMissing handles the deletion of missing files and triggers necessary cleanup operations func (s *maintenanceService) deleteMissing(ctx context.Context, ids []string) error { // Track affected album IDs before deletion for refresh @@ -52,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) @@ -68,7 +162,8 @@ func (s *maintenanceService) deleteMissing(ctx context.Context, ids []string) er return err } - // Refresh statistics in background + // Refresh statistics in background. album/artist play count aggregates are not recalculated + // here; they are refreshed by the next scan. s.refreshStatsAsync(ctx, affectedAlbumIDs) return nil @@ -97,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 { @@ -115,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", }) @@ -148,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 @@ -170,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 { @@ -198,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 09b442438..4ffc098d9 100644 --- a/core/maintenance_test.go +++ b/core/maintenance_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "sync" + "time" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" @@ -250,6 +251,167 @@ var _ = Describe("Maintenance", func() { }) }) }) + + Describe("RemapMissingFile", func() { + It("relocates the missing file's identity onto the target and runs GC", func() { + created := time.Now().Add(-30 * 24 * time.Hour) + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", Path: "old/song.mp3", AlbumID: "album1", CreatedAt: created, Missing: true}, + {ID: "t1", Path: "new/song.mp3", AlbumID: "album1", Missing: false}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) + + 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(ctx, "t1") + Expect(exists).To(BeFalse()) // discarded row removed + Expect(ds.GCCalled).To(BeTrue()) + }) + + It("moves the target's annotations, bookmarks and playlist entries onto the surviving id", func() { + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", AlbumID: "album1", Missing: true}, + {ID: "t1", AlbumID: "album1", Missing: false}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) + + Expect(mfRepo.ReassignReferencesCalls).To(HaveKeyWithValue("t1", "m1")) + }) + + It("reassigns album annotations when the old album is emptied", func() { + albumRepo := ds.MockedAlbum.(*extendedAlbumRepo) + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", AlbumID: "album1", Missing: true}, + {ID: "t1", AlbumID: "album2", Missing: false}, + }) + + mfRepo.SetCountAll(0) + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) + + Expect(albumRepo.ReassignAnnotationCalls).To(HaveKeyWithValue("album1", "album2")) + }) + + It("does not reassign album annotations when the old album is not emptied", func() { + albumRepo := ds.MockedAlbum.(*extendedAlbumRepo) + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", AlbumID: "album1", Missing: true}, + {ID: "m2", AlbumID: "album1", Missing: false}, + {ID: "t1", AlbumID: "album2", Missing: false}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) + + Expect(albumRepo.ReassignAnnotationCalls).To(BeEmpty()) + }) + + It("does not reassign annotations when the album is unchanged", func() { + albumRepo := ds.MockedAlbum.(*extendedAlbumRepo) + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", AlbumID: "album1", Missing: true}, + {ID: "t1", AlbumID: "album1", Missing: false}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) + + Expect(albumRepo.ReassignAnnotationCalls).To(BeEmpty()) + }) + + It("returns ErrNotFound when the missing file does not exist", func() { + mfRepo.SetData(model.MediaFiles{{ID: "t1", Missing: false}}) + + Expect(service.RemapMissingFile(ctx, "nope", "t1")).To(MatchError(model.ErrNotFound)) + }) + + It("refuses to remap from a file that is not missing", func() { + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", Missing: false}, + {ID: "t1", Missing: false}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(MatchError(ErrNotMissing)) + }) + + It("refuses to remap a file onto itself", func() { + mfRepo.SetData(model.MediaFiles{{ID: "m1", Missing: true}}) + + Expect(service.RemapMissingFile(ctx, "m1", "m1")).To(MatchError(ErrSameFile)) + }) + + It("refuses to remap onto a target that is itself missing", func() { + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", Missing: true}, + {ID: "t1", Missing: true}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(MatchError(ErrTargetMissing)) + }) + + It("refreshes artist and album stats right after the remap", func() { + artistRepo := ds.MockedArtist.(*extendedArtistRepo) + albumRepo := ds.MockedAlbum.(*extendedAlbumRepo) + albumRepo.SetData(model.Albums{ + {ID: "album1", Name: "Old Album", SongCount: 2, Size: 1100, Duration: 110}, + {ID: "album2", Name: "New Album", SongCount: 1, Size: 2000, Duration: 200}, + }) + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", Path: "old/1.mp3", Album: "Old Album", AlbumID: "album1", Missing: true, Size: 100, Duration: 10}, + {ID: "k1", Path: "old/2.mp3", Album: "Old Album", AlbumID: "album1", Missing: false, Size: 1000, Duration: 100}, + {ID: "t1", Path: "new/1.mp3", Album: "New Album", AlbumID: "album2", Missing: false, Size: 2000, Duration: 200}, + }) + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(Succeed()) + + 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(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(ctx, "album2") + Expect(err).ToNot(HaveOccurred()) + Expect(newAlbum.SongCount).To(Equal(1)) + Expect(newAlbum.Size).To(Equal(int64(2000))) + Expect(newAlbum.Duration).To(BeNumerically("==", 200)) + }) + + It("returns an error if GC fails", func() { + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", AlbumID: "album1", Missing: true}, + {ID: "t1", AlbumID: "album1", Missing: false}, + }) + ds.GCError = errors.New("gc failed") + + Expect(service.RemapMissingFile(ctx, "m1", "t1")).To(MatchError(ContainSubstring("gc failed"))) + }) + + It("preserves the target's participants on the remapped track", func() { + participant := model.Participant{ + Artist: model.Artist{ID: "a1", Name: "Artist", OrderArtistName: "artist", MbzArtistID: "mbz-artist"}, + } + mfRepo.SetData(model.MediaFiles{ + {ID: "m1", AlbumID: "album1", Missing: true}, + {ID: "t1", AlbumID: "album2", Missing: false, Participants: model.Participants{ + model.RoleArtist: model.ParticipantList{participant}, + }}, + }) + + 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(ctx, "m1") + Expect(err).ToNot(HaveOccurred()) + Expect(got.Participants).To(HaveKeyWithValue(model.RoleArtist, model.ParticipantList{participant})) + }) + }) }) // Test helper to create a mock DataStore with controllable behavior @@ -285,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 { @@ -308,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 @@ -328,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 { @@ -345,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 @@ -354,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..1d2ff5df4 100644 --- a/core/metrics/insights.go +++ b/core/metrics/insights.go @@ -10,6 +10,7 @@ import ( "path/filepath" "runtime" "runtime/debug" + "slices" "strings" "sync" "sync/atomic" @@ -47,11 +48,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 +88,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 +246,45 @@ 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() + libs, err := c.ds.Library().GetAll(ctx) if err != nil { - log.Trace(ctx, "Error reading libraries count", err) + log.Trace(ctx, "Error reading libraries", err) } - data.Library.ActiveUsers, err = c.ds.User(ctx).CountAll(model.QueryOptions{ + data.Library.Libraries = int64(len(libs)) + if slices.ContainsFunc(libs, func(lib model.Library) bool { return lib.PIDAlbum != "" || lib.PIDTrack != "" }) { + data.Config.HasCustomPID = true + } + 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 +302,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 +329,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 963914514..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,18 +34,17 @@ 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 != "" { - plr, err = p.ds.Player(ctx).Get(playerID) - if err == nil && plr.Client != client { + 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 { @@ -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,17 +80,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) - 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(ctx).Get(plr.TranscodingId) + if plr.TranscodingId == "" { + return plr, nil, nil } + 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 90c265fcc..e452c52ba 100644 --- a/core/players_test.go +++ b/core/players_test.go @@ -61,8 +61,19 @@ var _ = Describe("Players", func() { Expect(trc).To(BeNil()) }) + It("does not reuse another user's player by ID", func() { + plr := &model.Player{ID: "123", Name: "A Player", Client: "client", UserId: "otheruser", UserAgent: "Pixel", TranscodingId: "1"} + repo.add(plr) + p, trc, err := players.Register(ctx, "123", "client", "chrome", "1.2.3.4") + Expect(err).ToNot(HaveOccurred()) + Expect(p.ID).ToNot(Equal("123")) + Expect(p.UserId).To(Equal("userid")) + Expect(repo.lastSaved).To(Equal(p)) + Expect(trc).To(BeNil()) + }) + It("finds players by ID", func() { - plr := &model.Player{ID: "123", Name: "A Player", Client: "client", LastSeen: time.Time{}} + plr := &model.Player{ID: "123", Name: "A Player", Client: "client", UserId: "userid", LastSeen: time.Time{}} repo.add(plr) p, trc, err := players.Register(ctx, "123", "client", "chrome", "1.2.3.4") Expect(err).ToNot(HaveOccurred()) @@ -93,7 +104,7 @@ var _ = Describe("Players", func() { }) It("finds player by ID and return its transcoding", func() { - plr := &model.Player{ID: "123", Name: "A Player", Client: "client", LastSeen: time.Time{}, TranscodingId: "1"} + plr := &model.Player{ID: "123", Name: "A Player", Client: "client", UserId: "userid", LastSeen: time.Time{}, TranscodingId: "1"} repo.add(plr) p, trc, err := players.Register(ctx, "123", "client", "chrome", "1.2.3.4") Expect(err).ToNot(HaveOccurred()) @@ -103,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"}) @@ -119,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 { @@ -134,14 +182,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 @@ -150,7 +198,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..b5991b095 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,17 +74,17 @@ 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 } - matcher := newLibraryMatcher(libs) - lib, ok := matcher.findLibrary(dir) + matcher := model.NewLibraryMatcher(libs) + lib, ok := matcher.FindLibrary(dir) if !ok { 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 445561266..25960f0fe 100644 --- a/core/playlists/import_test.go +++ b/core/playlists/import_test.go @@ -236,6 +236,24 @@ var _ = Describe("Playlists - Import", func() { Expect(pls.ExternalImageURL).To(BeEmpty()) }) + It("rejects #EXTALBUMARTURL pointing at a non-image file inside the library", func() { + tmpDir := GinkgoT().TempDir() + Expect(os.WriteFile(filepath.Join(tmpDir, "config.ini"), []byte("password=secret"), 0600)).To(Succeed()) + + m3u := "#EXTALBUMARTURL:config.ini\ntest.mp3\n" + plsFile := filepath.Join(tmpDir, "test.m3u") + Expect(os.WriteFile(plsFile, []byte(m3u), 0600)).To(Succeed()) + + mockLibRepo.SetData([]model.Library{{ID: 1, Path: tmpDir}}) + ds.MockedMediaFile = &mockedMediaFileFromListRepo{data: []string{"test.mp3"}} + ps = playlists.NewPlaylists(ds, artwork.NewUploader(ds)) + + plsFolder := &model.Folder{ID: "1", LibraryID: 1, LibraryPath: tmpDir, Path: "", Name: ""} + pls, err := ps.ImportFromFolder(ctx, plsFolder, "test.m3u") + Expect(err).ToNot(HaveOccurred()) + Expect(pls.ExternalImageURL).To(BeEmpty()) + }) + It("ignores HTTP #EXTALBUMARTURL when EnableM3UExternalAlbumArt is false", func() { conf.Server.EnableM3UExternalAlbumArt = false @@ -1011,6 +1029,19 @@ var _ = Describe("Playlists - Import", func() { Expect(pls.ExternalImageURL).To(BeEmpty()) }) + DescribeTable("restricts a local #EXTALBUMARTURL to the owner's libraries", + func(imageURL, expected string) { + ctx = request.WithUser(ctx, model.User{ID: "123", Libraries: model.Libraries{{ID: 1, Path: "/music"}}}) + repo.data = []string{"tests/test.mp3"} + m3u := "#EXTALBUMARTURL:" + imageURL + "\n/music/tests/test.mp3\n" + pls, err := ps.ImportM3U(ctx, strings.NewReader(m3u)) + Expect(err).ToNot(HaveOccurred()) + Expect(pls.ExternalImageURL).To(Equal(expected)) + }, + Entry("accepts a library the owner can access", "file:///music/cover.jpg", filepath.Clean("/music/cover.jpg")), + Entry("ignores a library the owner cannot access", "file:///new/cover.jpg", ""), + ) + // Fullwidth characters (e.g., ABCD) are not handled by SQLite's NOCASE collation, // so we need exact matching for non-ASCII characters. It("matches fullwidth characters exactly (SQLite NOCASE limitation)", func() { @@ -1141,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 @@ -1181,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 { @@ -1216,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 286f2e420..ab95b8850 100644 --- a/core/playlists/parse_m3u.go +++ b/core/playlists/parse_m3u.go @@ -1,25 +1,24 @@ package playlists import ( - "cmp" "context" "fmt" "io" "net/url" "path/filepath" - "slices" "strings" "time" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/utils/slice" "golang.org/x/text/unicode/norm" ) 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 @@ -36,7 +35,8 @@ func (s *playlists) parseM3U(ctx context.Context, pls *model.Playlist, folder *m continue } if after, ok := strings.CutPrefix(line, "#EXTALBUMARTURL:"); ok { - pls.ExternalImageURL = resolveImageURL(after, folder, resolver.matcher) + owner, _ := request.UserFrom(ctx) + pls.ExternalImageURL = resolveImageURL(after, folder, resolver.matcher, owner) continue } // Skip empty lines and extended info @@ -94,7 +94,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 @@ -154,70 +154,18 @@ func (r pathResolution) ToQualifiedString() (string, error) { return fmt.Sprintf("%d:%s", r.libraryID, filepath.ToSlash(relativePath)), nil } -// libraryMatcher holds sorted libraries with cleaned paths for efficient path matching. -type libraryMatcher struct { - libraries model.Libraries - cleanedPaths []string -} - -// findLibraryForPath finds which library contains the given absolute path. -// Returns library ID and path, or 0 and empty string if not found. -func (lm *libraryMatcher) findLibraryForPath(absolutePath string) (int, string) { - lib, ok := lm.findLibrary(absolutePath) - if !ok { - return 0, "" - } - return lib.ID, filepath.Clean(lib.Path) -} - -// findLibrary checks if the absolute path is under any of the library paths. -func (lm *libraryMatcher) findLibrary(absolutePath string) (model.Library, bool) { - // Check sorted libraries (longest path first) to find the best match - for i, cleanLibPath := range lm.cleanedPaths { - // Check if absolutePath is under this library path - if strings.HasPrefix(absolutePath, cleanLibPath) { - // Ensure it's a proper path boundary (not just a prefix) - if len(absolutePath) == len(cleanLibPath) || absolutePath[len(cleanLibPath)] == filepath.Separator { - return lm.libraries[i], true - } - } - } - return model.Library{}, false -} - -// newLibraryMatcher creates a libraryMatcher with libraries sorted by path length (longest first). -// This ensures correct matching when library paths are prefixes of each other. -// Example: /music-classical must be checked before /music -// Otherwise, /music-classical/track.mp3 would match /music instead of /music-classical -func newLibraryMatcher(libs model.Libraries) *libraryMatcher { - // Sort libraries by path length (descending) to ensure longest paths match first. - slices.SortFunc(libs, func(i, j model.Library) int { - return cmp.Compare(len(j.Path), len(i.Path)) // Reverse order for descending - }) - - // Pre-clean all library paths once for efficient matching - cleanedPaths := make([]string, len(libs)) - for i, lib := range libs { - cleanedPaths[i] = filepath.Clean(lib.Path) - } - return &libraryMatcher{ - libraries: libs, - cleanedPaths: cleanedPaths, - } -} - // pathResolver handles path resolution logic for playlist imports. type pathResolver struct { - matcher *libraryMatcher + matcher *model.LibraryMatcher } // 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 } - matcher := newLibraryMatcher(libs) + matcher := model.NewLibraryMatcher(libs) return &pathResolver{matcher: matcher}, nil } @@ -244,14 +192,14 @@ func (r *pathResolver) resolvePath(line string, folder *model.Folder) pathResolu // a pathResolution with the library information. Returns an invalid resolution if // the path is not found in any library. func (r *pathResolver) findInLibraries(absolutePath string) pathResolution { - libID, libPath := r.matcher.findLibraryForPath(absolutePath) - if libID == 0 { + lib, ok := r.matcher.FindLibrary(absolutePath) + if !ok { return pathResolution{valid: false} } return pathResolution{ absolutePath: absolutePath, - libraryPath: libPath, - libraryID: libID, + libraryPath: filepath.Clean(lib.Path), + libraryID: lib.ID, valid: true, } } @@ -286,7 +234,7 @@ func (r *pathResolver) resolvePaths(ctx context.Context, folder *model.Folder, l // HTTP(S) URLs are stored as-is (gated by EnableM3UExternalAlbumArt). // Local paths (file://, absolute, or relative) are resolved to an absolute path // and validated against known library boundaries via matcher. -func resolveImageURL(value string, folder *model.Folder, matcher *libraryMatcher) string { +func resolveImageURL(value string, folder *model.Folder, matcher *model.LibraryMatcher, owner model.User) string { value = strings.TrimSpace(value) if value == "" { return "" @@ -302,12 +250,13 @@ func resolveImageURL(value string, folder *model.Folder, matcher *libraryMatcher // Resolve to local absolute path localPath, ok := resolveLocalPath(value, folder) - if !ok { + if !ok || !model.IsImageFile(localPath) { return "" } - // Validate path is within a known library - if libID, _ := matcher.findLibraryForPath(localPath); libID == 0 { + lib, ok := matcher.FindLibrary(localPath) + // A playlist without a folder (API upload, or CLI import from outside all libraries) may only use the owner's libraries. + if !ok || (folder == nil && !owner.HasLibraryAccess(lib.ID)) { return "" } return localPath diff --git a/core/playlists/parse_m3u_test.go b/core/playlists/parse_m3u_test.go index d7fd5e001..ced6c2b16 100644 --- a/core/playlists/parse_m3u_test.go +++ b/core/playlists/parse_m3u_test.go @@ -9,187 +9,6 @@ import ( . "github.com/onsi/gomega" ) -var _ = Describe("libraryMatcher", func() { - var ds *tests.MockDataStore - var mockLibRepo *tests.MockLibraryRepo - ctx := context.Background() - - BeforeEach(func() { - tests.SkipOnWindows("path separator bug (#TBD-path-sep-playlists)") - mockLibRepo = &tests.MockLibraryRepo{} - ds = &tests.MockDataStore{ - MockedLibrary: mockLibRepo, - } - }) - - // Helper function to create a libraryMatcher from the mock datastore - createMatcher := func(ds model.DataStore) *libraryMatcher { - libs, err := ds.Library(ctx).GetAll() - Expect(err).ToNot(HaveOccurred()) - return newLibraryMatcher(libs) - } - - Describe("Longest library path matching", func() { - It("matches the longest library path when multiple libraries share a prefix", func() { - // Setup libraries with prefix conflicts - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/music"}, - {ID: 2, Path: "/music-classical"}, - {ID: 3, Path: "/music-classical/opera"}, - }) - - matcher := createMatcher(ds) - - // Test that longest path matches first and returns correct library ID - testCases := []struct { - path string - expectedLibID int - expectedLibPath string - }{ - {"/music-classical/opera/track.mp3", 3, "/music-classical/opera"}, - {"/music-classical/track.mp3", 2, "/music-classical"}, - {"/music/track.mp3", 1, "/music"}, - {"/music-classical/opera/subdir/file.mp3", 3, "/music-classical/opera"}, - } - - for _, tc := range testCases { - libID, libPath := matcher.findLibraryForPath(tc.path) - Expect(libID).To(Equal(tc.expectedLibID), "Path %s should match library ID %d, but got %d", tc.path, tc.expectedLibID, libID) - Expect(libPath).To(Equal(tc.expectedLibPath), "Path %s should match library path %s, but got %s", tc.path, tc.expectedLibPath, libPath) - } - }) - - It("handles libraries with similar prefixes but different structures", func() { - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/home/user/music"}, - {ID: 2, Path: "/home/user/music-backup"}, - }) - - matcher := createMatcher(ds) - - // Test that music-backup library is matched correctly - libID, libPath := matcher.findLibraryForPath("/home/user/music-backup/track.mp3") - Expect(libID).To(Equal(2)) - Expect(libPath).To(Equal("/home/user/music-backup")) - - // Test that music library is still matched correctly - libID, libPath = matcher.findLibraryForPath("/home/user/music/track.mp3") - Expect(libID).To(Equal(1)) - Expect(libPath).To(Equal("/home/user/music")) - }) - - It("matches path that is exactly the library root", func() { - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/music"}, - {ID: 2, Path: "/music-classical"}, - }) - - matcher := createMatcher(ds) - - // Exact library path should match - libID, libPath := matcher.findLibraryForPath("/music-classical") - Expect(libID).To(Equal(2)) - Expect(libPath).To(Equal("/music-classical")) - }) - - It("handles complex nested library structures", func() { - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/media"}, - {ID: 2, Path: "/media/audio"}, - {ID: 3, Path: "/media/audio/classical"}, - {ID: 4, Path: "/media/audio/classical/baroque"}, - }) - - matcher := createMatcher(ds) - - testCases := []struct { - path string - expectedLibID int - expectedLibPath string - }{ - {"/media/audio/classical/baroque/bach/track.mp3", 4, "/media/audio/classical/baroque"}, - {"/media/audio/classical/mozart/track.mp3", 3, "/media/audio/classical"}, - {"/media/audio/rock/track.mp3", 2, "/media/audio"}, - {"/media/video/movie.mp4", 1, "/media"}, - } - - for _, tc := range testCases { - libID, libPath := matcher.findLibraryForPath(tc.path) - Expect(libID).To(Equal(tc.expectedLibID), "Path %s should match library ID %d", tc.path, tc.expectedLibID) - Expect(libPath).To(Equal(tc.expectedLibPath), "Path %s should match library path %s", tc.path, tc.expectedLibPath) - } - }) - }) - - Describe("Edge cases", func() { - It("handles empty library list", func() { - mockLibRepo.SetData([]model.Library{}) - - matcher := createMatcher(ds) - Expect(matcher).ToNot(BeNil()) - - // Should not match anything - libID, libPath := matcher.findLibraryForPath("/music/track.mp3") - Expect(libID).To(Equal(0)) - Expect(libPath).To(BeEmpty()) - }) - - It("handles single library", func() { - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/music"}, - }) - - matcher := createMatcher(ds) - - libID, libPath := matcher.findLibraryForPath("/music/track.mp3") - Expect(libID).To(Equal(1)) - Expect(libPath).To(Equal("/music")) - }) - - It("handles libraries with special characters in paths", func() { - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/music[test]"}, - {ID: 2, Path: "/music(backup)"}, - }) - - matcher := createMatcher(ds) - Expect(matcher).ToNot(BeNil()) - - // Special characters should match literally - libID, libPath := matcher.findLibraryForPath("/music[test]/track.mp3") - Expect(libID).To(Equal(1)) - Expect(libPath).To(Equal("/music[test]")) - }) - }) - - Describe("Path matching order", func() { - It("ensures longest paths match first", func() { - mockLibRepo.SetData([]model.Library{ - {ID: 1, Path: "/a"}, - {ID: 2, Path: "/ab"}, - {ID: 3, Path: "/abc"}, - }) - - matcher := createMatcher(ds) - - // Verify that longer paths match correctly (not cut off by shorter prefix) - testCases := []struct { - path string - expectedLibID int - }{ - {"/abc/file.mp3", 3}, - {"/ab/file.mp3", 2}, - {"/a/file.mp3", 1}, - } - - for _, tc := range testCases { - libID, _ := matcher.findLibraryForPath(tc.path) - Expect(libID).To(Equal(tc.expectedLibID), "Path %s should match library ID %d", tc.path, tc.expectedLibID) - } - }) - }) -}) - var _ = Describe("pathResolver", func() { var ds *tests.MockDataStore var mockLibRepo *tests.MockLibraryRepo diff --git a/core/playlists/playlists.go b/core/playlists/playlists.go index 9ae8f09cb..c9bc03b97 100644 --- a/core/playlists/playlists.go +++ b/core/playlists/playlists.go @@ -32,6 +32,7 @@ type Playlists interface { // Track management AddTracks(ctx context.Context, playlistID string, ids []string) (int, error) + InsertTracks(ctx context.Context, playlistID string, ids []string, pos int) (int, error) AddAlbums(ctx context.Context, playlistID string, albumIds []string) (int, error) AddArtists(ctx context.Context, playlistID string, artistIds []string) (int, error) AddDiscs(ctx context.Context, playlistID string, discs []model.DiscID) (int, error) @@ -48,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. @@ -63,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 { @@ -85,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 } @@ -126,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 } @@ -144,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 }) @@ -164,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, @@ -182,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 } } @@ -204,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 } } @@ -220,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 } @@ -256,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 --- @@ -265,28 +269,44 @@ 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. +func (s *playlists) InsertTracks(ctx context.Context, playlistID string, ids []string, pos int) (int, error) { + if _, err := s.checkTracksEditable(ctx, playlistID); err != nil { + return 0, err + } + var count int + // Immediate: a deferred tx that reads before shifting fails at once with SQLITE_BUSY under + // concurrent writers instead of waiting for the lock. + err := s.ds.WithTxImmediate(func(tx model.DataStore) error { + var err error + count, err = tx.Playlist().Tracks(ctx, playlistID, false).Insert(ctx, ids, pos) + return err + }) + return count, err } func (s *playlists) AddAlbums(ctx context.Context, playlistID string, albumIds []string) (int, error) { 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 { @@ -294,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...) }) } @@ -302,8 +322,8 @@ func (s *playlists) ReorderTrack(ctx context.Context, playlistID string, pos int if _, err := s.checkTracksEditable(ctx, playlistID); err != nil { return err } - return s.ds.WithTx(func(tx model.DataStore) error { - return tx.Playlist(ctx).Tracks(playlistID, false).Reorder(pos, newPos) + return s.ds.WithTxImmediate(func(tx model.DataStore) error { + return tx.Playlist().Tracks(ctx, playlistID, false).Reorder(ctx, pos, newPos) }) } @@ -322,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) @@ -340,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 dd182b3f1..ec44de3bc 100644 --- a/core/playlists/playlists_test.go +++ b/core/playlists/playlists_test.go @@ -324,6 +324,43 @@ var _ = Describe("Playlists", func() { }) }) + Describe("InsertTracks", func() { + var mockTracks *tests.MockPlaylistTrackRepo + + BeforeEach(func() { + mockTracks = &tests.MockPlaylistTrackRepo{AddCount: 2} + mockPlsRepo.Data = map[string]*model.Playlist{ + "pls-1": {ID: "pls-1", Name: "My Playlist", OwnerID: "user-1"}, + "pls-smart": {ID: "pls-smart", Name: "Smart", OwnerID: "user-1", + Rules: &criteria.Criteria{Expression: criteria.Contains{"title": "test"}}}, + } + mockPlsRepo.TracksRepo = mockTracks + ps = playlists.NewPlaylists(ds, artwork.NewUploader(ds)) + }) + + It("inserts the tracks at the given position for the owner", func() { + ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) + count, err := ps.InsertTracks(ctx, "pls-1", []string{"song-1", "song-2"}, 3) + Expect(err).ToNot(HaveOccurred()) + Expect(count).To(Equal(2)) + Expect(mockTracks.AddedIds).To(Equal([]string{"song-1", "song-2"})) + Expect(mockTracks.InsertPos).To(Equal(3)) + }) + + It("denies non-owner, non-admin", func() { + ctx = request.WithUser(ctx, model.User{ID: "other-user", IsAdmin: false}) + _, err := ps.InsertTracks(ctx, "pls-1", []string{"song-1"}, 1) + Expect(err).To(MatchError(model.ErrNotAuthorized)) + Expect(mockTracks.AddedIds).To(BeEmpty()) + }) + + It("denies editing smart playlists", func() { + ctx = request.WithUser(ctx, model.User{ID: "user-1", IsAdmin: false}) + _, err := ps.InsertTracks(ctx, "pls-smart", []string{"song-1"}, 1) + Expect(err).To(MatchError(model.ErrPlaylistNotEditable)) + }) + }) + Describe("ReorderTrack", func() { var mockTracks *tests.MockPlaylistTrackRepo @@ -456,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 ca7cfb0cc..f8e8a9d41 100644 --- a/core/playlists/rest_adapter.go +++ b/core/playlists/rest_adapter.go @@ -2,7 +2,6 @@ package playlists import ( "context" - "errors" "reflect" "strings" @@ -15,50 +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 { - err := r.service.Delete(r.ctx, id) - switch { - case errors.Is(err, model.ErrNotFound): - return rest.ErrNotFound - case errors.Is(err, model.ErrNotAuthorized): - return rest.ErrPermissionDenied - default: - return err +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. @@ -72,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 } @@ -92,14 +84,7 @@ func (s *playlists) savePlaylist(ctx context.Context, pls *model.Playlist) (stri func (s *playlists) updatePlaylistEntity(ctx context.Context, id string, entity *model.Playlist, cols ...string) error { current, err := s.checkWritable(ctx, id) if err != nil { - switch { - case errors.Is(err, model.ErrNotFound): - return rest.ErrNotFound - case errors.Is(err, model.ErrNotAuthorized): - return rest.ErrPermissionDenied - default: - return err - } + return err } sent := sentFields(cols) @@ -172,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/quickconnect/quickconnect.go b/core/quickconnect/quickconnect.go new file mode 100644 index 000000000..9b69bc904 --- /dev/null +++ b/core/quickconnect/quickconnect.go @@ -0,0 +1,203 @@ +// Package quickconnect implements Jellyfin-style Quick Connect: a new client shows a short code, and +// an already signed-in user approves it, so the client can sign in without a password. +package quickconnect + +import ( + "crypto/rand" + "encoding/hex" + "errors" + "fmt" + "strings" + "sync" + "time" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/utils/random" + "github.com/navidrome/navidrome/utils/singleton" +) + +const timeout = 10 * time.Minute + +// Initiate is unauthenticated, so the number of live requests must be bounded. +const maxPending = 1000 + +var ( + ErrAlreadyAuthorized = errors.New("quick connect request already authorized") + ErrTooManyRequests = errors.New("too many pending quick connect requests") +) + +type Device struct { + ID string + Name string + App string + AppVersion string +} + +type Request struct { + Device Device + Secret string + Code string + DateAdded time.Time + UserID string +} + +func (r Request) Authorized() bool { + return r.UserID != "" +} + +type QuickConnect interface { + Initiate(device Device) (Request, error) + Status(secret string) (Request, error) + // Lookup returns a request still waiting for approval. + Lookup(code string) (Request, error) + Authorize(code, userID string) (Request, error) + // Redeem returns the id of the user who approved the request. It succeeds only once per secret. + Redeem(secret string) (string, error) +} + +type entry struct { + Request + expiresAt time.Time +} + +type quickConnect struct { + mu sync.Mutex + bySecret map[string]*entry + byCode map[string]*entry +} + +// GetInstance returns the shared store: the Jellyfin and native API routers must see the same requests. +func GetInstance() QuickConnect { + return singleton.GetInstance(newStore) +} + +func New() QuickConnect { + return newStore() +} + +func newStore() *quickConnect { + return &quickConnect{ + bySecret: map[string]*entry{}, + byCode: map[string]*entry{}, + } +} + +// Enabled reports whether users can approve codes from the web UI, which needs the Jellyfin API. +func Enabled() bool { + return conf.Server.Jellyfin.Enabled && conf.Server.Jellyfin.QuickConnect +} + +func (qc *quickConnect) Initiate(device Device) (Request, error) { + qc.mu.Lock() + defer qc.mu.Unlock() + qc.expire() + if len(qc.bySecret) >= maxPending { + return Request{}, ErrTooManyRequests + } + // The fields may be substrings of a much larger header; copy them so the header isn't retained. + device = Device{ID: strings.Clone(device.ID), Name: strings.Clone(device.Name), + App: strings.Clone(device.App), AppVersion: strings.Clone(device.AppVersion)} + now := time.Now() + e := &entry{ + Request: Request{Device: device, Secret: newSecret(), Code: qc.newCode(), DateAdded: now}, + expiresAt: now.Add(timeout), + } + qc.bySecret[e.Secret] = e + qc.byCode[e.Code] = e + return e.Request, nil +} + +func (qc *quickConnect) Status(secret string) (Request, error) { + qc.mu.Lock() + defer qc.mu.Unlock() + e, err := qc.find(qc.bySecret, secret) + if err != nil { + return Request{}, err + } + return e.Request, nil +} + +func (qc *quickConnect) Lookup(code string) (Request, error) { + qc.mu.Lock() + defer qc.mu.Unlock() + e, err := qc.findPending(code) + if err != nil { + return Request{}, err + } + return e.Request, nil +} + +func (qc *quickConnect) Authorize(code, userID string) (Request, error) { + qc.mu.Lock() + defer qc.mu.Unlock() + e, err := qc.findPending(code) + if err != nil { + return Request{}, err + } + e.UserID = userID + e.expiresAt = time.Now().Add(timeout) + return e.Request, nil +} + +func (qc *quickConnect) Redeem(secret string) (string, error) { + qc.mu.Lock() + defer qc.mu.Unlock() + e, err := qc.find(qc.bySecret, secret) + if err != nil || !e.Authorized() { + return "", model.ErrNotFound + } + qc.remove(e) + return e.UserID, nil +} + +func (qc *quickConnect) find(index map[string]*entry, key string) (*entry, error) { + qc.expire() + e, ok := index[key] + if !ok { + return nil, model.ErrNotFound + } + return e, nil +} + +func (qc *quickConnect) findPending(code string) (*entry, error) { + e, err := qc.find(qc.byCode, normalizeCode(code)) + if err == nil && e.Authorized() { + return nil, ErrAlreadyAuthorized + } + return e, err +} + +func (qc *quickConnect) expire() { + now := time.Now() + for _, e := range qc.bySecret { + if !now.Before(e.expiresAt) { + qc.remove(e) + } + } +} + +func (qc *quickConnect) remove(e *entry) { + delete(qc.bySecret, e.Secret) + delete(qc.byCode, e.Code) +} + +func (qc *quickConnect) newCode() string { + for { + code := fmt.Sprintf("%06d", random.Int64N(900000)+100000) + if _, taken := qc.byCode[code]; !taken { + return code + } + } +} + +func newSecret() string { + b := make([]byte, 32) + _, _ = rand.Read(b) + return hex.EncodeToString(b) +} + +// Users often type the code in groups ("123 456"). +func normalizeCode(code string) string { + return strings.Join(strings.Fields(code), "") +} diff --git a/core/quickconnect/quickconnect_suite_test.go b/core/quickconnect/quickconnect_suite_test.go new file mode 100644 index 000000000..029f4e333 --- /dev/null +++ b/core/quickconnect/quickconnect_suite_test.go @@ -0,0 +1,17 @@ +package quickconnect + +import ( + "testing" + + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestQuickConnect(t *testing.T) { + tests.Init(t, false) + log.SetLevel(log.LevelFatal) + RegisterFailHandler(Fail) + RunSpecs(t, "QuickConnect Suite") +} diff --git a/core/quickconnect/quickconnect_test.go b/core/quickconnect/quickconnect_test.go new file mode 100644 index 000000000..b03186807 --- /dev/null +++ b/core/quickconnect/quickconnect_test.go @@ -0,0 +1,190 @@ +package quickconnect + +import ( + "testing" + "testing/synctest" + "time" + + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var device = Device{ID: "dev-1", Name: "Pixel 7", App: "Finamp", AppVersion: "1.0.0"} + +var _ = Describe("QuickConnect", func() { + var qc QuickConnect + + BeforeEach(func() { + qc = New() + }) + + Describe("Initiate", func() { + It("creates a pending request with a 6-digit code and a random secret", func() { + req, err := qc.Initiate(device) + Expect(err).ToNot(HaveOccurred()) + Expect(req.Code).To(MatchRegexp(`^[1-9]\d{5}$`)) + Expect(req.Secret).To(MatchRegexp(`^[0-9a-f]{64}$`)) + Expect(req.Device).To(Equal(device)) + Expect(req.DateAdded).ToNot(BeZero()) + Expect(req.Authorized()).To(BeFalse()) + }) + + It("never gives two live requests the same code or secret", func() { + codes := map[string]bool{} + secrets := map[string]bool{} + for range 500 { + req, err := qc.Initiate(device) + Expect(err).ToNot(HaveOccurred()) + codes[req.Code] = true + secrets[req.Secret] = true + } + Expect(codes).To(HaveLen(500)) + Expect(secrets).To(HaveLen(500)) + }) + + It("refuses new requests when too many are pending", func() { + for range maxPending { + _, err := qc.Initiate(device) + Expect(err).ToNot(HaveOccurred()) + } + _, err := qc.Initiate(device) + Expect(err).To(MatchError(ErrTooManyRequests)) + }) + }) + + Describe("Status", func() { + It("returns the request for a known secret", func() { + req, _ := qc.Initiate(device) + got, err := qc.Status(req.Secret) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(req)) + }) + + It("returns ErrNotFound for an unknown secret", func() { + _, err := qc.Status("nope") + Expect(err).To(MatchError(model.ErrNotFound)) + }) + }) + + Describe("Lookup", func() { + It("finds a request by code, ignoring spaces", func() { + req, _ := qc.Initiate(device) + got, err := qc.Lookup(req.Code[:3] + " " + req.Code[3:]) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(req)) + }) + + It("returns ErrNotFound for an unknown code", func() { + _, err := qc.Lookup("000000") + Expect(err).To(MatchError(model.ErrNotFound)) + }) + + It("refuses a request that is already authorized", func() { + req, _ := qc.Initiate(device) + _, _ = qc.Authorize(req.Code, "user-1") + _, err := qc.Lookup(req.Code) + Expect(err).To(MatchError(ErrAlreadyAuthorized)) + }) + }) + + Describe("Authorize", func() { + It("marks the request as authorized by the user", func() { + req, _ := qc.Initiate(device) + got, err := qc.Authorize(" "+req.Code+" ", "user-1") + Expect(err).ToNot(HaveOccurred()) + Expect(got.Authorized()).To(BeTrue()) + Expect(got.UserID).To(Equal("user-1")) + + status, _ := qc.Status(req.Secret) + Expect(status.Authorized()).To(BeTrue()) + }) + + It("refuses a request that is already authorized", func() { + req, _ := qc.Initiate(device) + _, _ = qc.Authorize(req.Code, "user-1") + _, err := qc.Authorize(req.Code, "user-2") + Expect(err).To(MatchError(ErrAlreadyAuthorized)) + }) + + It("returns ErrNotFound for an unknown code", func() { + _, err := qc.Authorize("000000", "user-1") + Expect(err).To(MatchError(model.ErrNotFound)) + }) + }) + + Describe("Redeem", func() { + It("returns the approving user once, then forgets the request", func() { + req, _ := qc.Initiate(device) + _, _ = qc.Authorize(req.Code, "user-1") + + userID, err := qc.Redeem(req.Secret) + Expect(err).ToNot(HaveOccurred()) + Expect(userID).To(Equal("user-1")) + + _, err = qc.Redeem(req.Secret) + Expect(err).To(MatchError(model.ErrNotFound)) + _, err = qc.Status(req.Secret) + Expect(err).To(MatchError(model.ErrNotFound)) + _, err = qc.Lookup(req.Code) + Expect(err).To(MatchError(model.ErrNotFound)) + }) + + It("refuses a request that is not authorized yet", func() { + req, _ := qc.Initiate(device) + _, err := qc.Redeem(req.Secret) + Expect(err).To(MatchError(model.ErrNotFound)) + _, err = qc.Status(req.Secret) + Expect(err).ToNot(HaveOccurred()) + }) + }) +}) + +// Plain tests: testing/synctest needs a *testing.T, which Ginkgo doesn't give. +func TestPendingRequestExpires(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + g := NewWithT(t) + qc := New() + req, _ := qc.Initiate(device) + + time.Sleep(timeout - time.Second) + _, err := qc.Status(req.Secret) + g.Expect(err).ToNot(HaveOccurred()) + + time.Sleep(2 * time.Second) + _, err = qc.Status(req.Secret) + g.Expect(err).To(MatchError(model.ErrNotFound)) + _, err = qc.Authorize(req.Code, "user-1") + g.Expect(err).To(MatchError(model.ErrNotFound)) + }) +} + +func TestAuthorizeExtendsExpiry(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + g := NewWithT(t) + qc := New() + req, _ := qc.Initiate(device) + + time.Sleep(timeout - time.Second) + _, err := qc.Authorize(req.Code, "user-1") + g.Expect(err).ToNot(HaveOccurred()) + + time.Sleep(timeout - time.Second) + userID, err := qc.Redeem(req.Secret) + g.Expect(err).ToNot(HaveOccurred()) + g.Expect(userID).To(Equal("user-1")) + }) +} + +func TestExpiredRequestsFreeCapacity(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + g := NewWithT(t) + qc := New() + for range maxPending { + _, _ = qc.Initiate(device) + } + time.Sleep(timeout + time.Second) + _, err := qc.Initiate(device) + g.Expect(err).ToNot(HaveOccurred()) + }) +} 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 63397f8f6..541a8e92b 100644 --- a/core/scrobbler/play_tracker.go +++ b/core/scrobbler/play_tracker.go @@ -17,6 +17,7 @@ import ( "github.com/navidrome/navidrome/server/events" "github.com/navidrome/navidrome/utils/cache" "github.com/navidrome/navidrome/utils/singleton" + "github.com/navidrome/navidrome/utils/slice" ) const ( @@ -67,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 } @@ -292,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 } @@ -327,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 } @@ -363,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 } @@ -408,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 } @@ -444,8 +445,14 @@ func (p *playTracker) ReportPlayback(ctx context.Context, params ReportPlaybackP return nil } -func (p *playTracker) GetNowPlaying(_ context.Context) ([]PlaybackSession, error) { +func (p *playTracker) GetNowPlaying(ctx context.Context) ([]PlaybackSession, error) { + // The cache is process-global, so it holds every user's playback, across all libraries. res := p.playMap.Values() + if user, ok := request.UserFrom(ctx); ok { + res = slice.Filter(res, func(s PlaybackSession) bool { + return user.HasLibraryAccess(s.MediaFile.LibraryID) + }) + } slices.SortFunc(res, func(a, b PlaybackSession) int { return b.Start.Compare(a.Start) }) @@ -470,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 @@ -497,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 }) @@ -531,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 0e768c3f4..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() { @@ -92,7 +92,7 @@ var _ = Describe("PlayTracker", func() { BeforeEach(func() { DeferCleanup(configtest.SetupConfig()) ctx = GinkgoT().Context() - ctx = request.WithUser(ctx, model.User{ID: "u-1"}) + ctx = request.WithUser(ctx, model.User{ID: "u-1", Libraries: model.Libraries{{ID: 1}}}) ctx = request.WithPlayer(ctx, model.Player{ScrobbleEnabled: true}) ds = &tests.MockDataStore{} fake = &fakeScrobbler{Authorized: true} @@ -108,6 +108,7 @@ var _ = Describe("PlayTracker", func() { track = model.MediaFile{ ID: "123", + LibraryID: 1, Title: "Track Title", Album: "Track Album", AlbumID: "al-1", @@ -118,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() { @@ -148,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{ @@ -174,6 +175,46 @@ var _ = Describe("PlayTracker", func() { Expect(playing[1].Username).To(Equal("user-1")) Expect(playing[1].MediaFile.ID).To(Equal("123")) }) + + It("hides sessions playing from libraries the caller cannot access", func() { + hidden := track + hidden.ID = "789" + hidden.LibraryID = 2 + _ = ds.MediaFile().Put(ctx, &hidden) + reporter := request.WithPlayer( + request.WithUser(GinkgoT().Context(), model.User{ID: "u-2", UserName: "user-2"}), + model.Player{ScrobbleEnabled: true}, + ) + _ = tracker.ReportPlayback(reporter, ReportPlaybackParams{ + MediaId: "789", PositionMs: 0, State: StatePlaying, PlaybackRate: 1.0, ClientId: "player-2", ClientName: "player-two", + }) + + playing, err := tracker.GetNowPlaying(ctx) + + Expect(err).ToNot(HaveOccurred()) + Expect(playing).To(BeEmpty(), "u-1 is granted library 1 only") + }) + + It("shows every session to an admin", func() { + hidden := track + hidden.ID = "789" + hidden.LibraryID = 2 + _ = ds.MediaFile().Put(ctx, &hidden) + reporter := request.WithPlayer( + request.WithUser(GinkgoT().Context(), model.User{ID: "u-2", UserName: "user-2"}), + model.Player{ScrobbleEnabled: true}, + ) + _ = tracker.ReportPlayback(reporter, ReportPlaybackParams{ + MediaId: "789", PositionMs: 0, State: StatePlaying, PlaybackRate: 1.0, ClientId: "player-2", ClientName: "player-two", + }) + + adminCtx := request.WithUser(GinkgoT().Context(), model.User{ID: "adm", IsAdmin: true}) + playing, err := tracker.GetNowPlaying(adminCtx) + + Expect(err).ToNot(HaveOccurred()) + Expect(playing).To(HaveLen(1)) + Expect(playing[0].MediaFile.ID).To(Equal("789")) + }) }) Describe("Expiration events", func() { @@ -319,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")) @@ -335,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)) }) }) @@ -347,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() { @@ -448,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 } @@ -591,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, @@ -618,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", @@ -719,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, @@ -917,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() { @@ -982,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 b2f32ba39..33c8996ef 100644 --- a/core/share.go +++ b/core/share.go @@ -2,14 +2,16 @@ package core import ( "context" + "fmt" + "slices" "strings" "time" "github.com/Masterminds/squirrel" - "github.com/deluan/rest" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" . "github.com/navidrome/navidrome/utils/gg" "github.com/navidrome/navidrome/utils/nanoid" "github.com/navidrome/navidrome/utils/slice" @@ -18,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 } @@ -44,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 } @@ -87,9 +80,13 @@ func (r *shareRepositoryWrapper) newId() (string, error) { } } -func (r *shareRepositoryWrapper) Save(entity any) (string, error) { - s := entity.(*model.Share) - id, err := r.newId() +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(ctx); ok { + s.UserID = user.ID + } + id, err := r.newId(ctx) if err != nil { return "", err } @@ -98,81 +95,96 @@ func (r *shareRepositoryWrapper) Save(entity any) (string, error) { s.ExpiresAt = new(time.Now().Add(conf.Server.DefaultShareExpiration)) } - firstId, _, _ := strings.Cut(s.ResourceIDs, ",") - v, err := model.GetEntityByID(r.ctx, r.ds, firstId) + s.ResourceType, err = r.resourceType(ctx, s.ResourceIDs) if err != nil { return "", err } - switch v.(type) { - case *model.Artist: - s.ResourceType = "artist" - s.Contents = r.contentsLabelFromArtist(s.ID, s.ResourceIDs) - case *model.Album: - s.ResourceType = "album" - s.Contents = r.contentsLabelFromAlbums(s.ID, s.ResourceIDs) - case *model.Playlist: - s.ResourceType = "playlist" - s.Contents = r.contentsLabelFromPlaylist(s.ID, s.ResourceIDs) - case *model.MediaFile: - s.ResourceType = "media_file" - s.Contents = r.contentsLabelFromMediaFiles(s.ID, s.ResourceIDs) - default: - log.Error(r.ctx, "Invalid Resource ID", "id", firstId) - return "", model.ErrNotFound + switch s.ResourceType { + case "artist": + s.Contents = r.contentsLabelFromArtist(ctx, s.ID, s.ResourceIDs) + case "album": + s.Contents = r.contentsLabelFromAlbums(ctx, s.ID, s.ResourceIDs) + case "playlist": + s.Contents = r.contentsLabelFromPlaylist(ctx, s.ID, s.ResourceIDs) + case "media_file": + 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) } -func (r *shareRepositoryWrapper) Update(id string, entity any, _ ...string) error { +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(ctx context.Context, resourceIDs string) (string, error) { + resourceType := "" + for _, id := range strings.Split(resourceIDs, ",") { + kind, err := model.GetEntityKindByID(ctx, r.ds, id) + if err != nil { + return "", err + } + if !slices.Contains(shareableKinds, kind) { + log.Error(ctx, "Invalid Resource ID", "id", id) + return "", model.ErrNotFound + } + if resourceType != "" && kind.String() != resourceType { + return "", fmt.Errorf("%w: share mixes %s and %s resources", model.ErrValidation, resourceType, kind) + } + resourceType = kind.String() + } + return resourceType, nil +} + +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 475d40ec9..8ce2bf270 100644 --- a/core/share_test.go +++ b/core/share_test.go @@ -5,6 +5,7 @@ import ( "github.com/deluan/rest" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -13,69 +14,90 @@ 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)) }) + 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.Repository().(rest.Persistable[model.Share]) + entity := &model.Share{Description: "test", ResourceIDs: "123", UserID: "victim-user"} + _, 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(ctx, entity) + Expect(err).To(MatchError(model.ErrNotFound)) + }) + + It("fails when the resource IDs are of mixed types", func() { + _ = ds.MediaFile().Put(ctx, &model.MediaFile{ID: "456", Title: "Example Media File"}) + entity := &model.Share{Description: "test", ResourceIDs: "123,456"} + _, 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/storage/local/watcher.go b/core/storage/local/watcher.go index 1b8a4e0c8..ef404d497 100644 --- a/core/storage/local/watcher.go +++ b/core/storage/local/watcher.go @@ -3,6 +3,7 @@ package local import ( "context" "errors" + "fmt" "path/filepath" "strings" @@ -18,22 +19,18 @@ func (s *localStorage) Start(ctx context.Context) (<-chan string, error) { return nil, errors.New("watcher already started") } input := make(chan notify.EventInfo, 500) - output := make(chan string, 500) + libPath := filepath.Join(s.u.Path, "...") + log.Debug(ctx, "Starting watcher", "lib", libPath) + if err := notify.Watch(libPath, input, WatchEvents); err != nil { + s.watching.Store(false) + return nil, fmt.Errorf("starting watcher on %s: %w", libPath, err) + } - started := make(chan struct{}) + output := make(chan string, 500) go func() { defer close(input) defer close(output) - - libPath := filepath.Join(s.u.Path, "...") - log.Debug(ctx, "Starting watcher", "lib", libPath) - err := notify.Watch(libPath, input, WatchEvents) - if err != nil { - log.Error("Error starting watcher", "lib", libPath, err) - return - } defer notify.Stop(input) - close(started) // signals the main goroutine we have started for { select { @@ -49,9 +46,5 @@ func (s *localStorage) Start(ctx context.Context) (<-chan string, error) { } } }() - select { - case <-started: - case <-ctx.Done(): - } return output, nil } diff --git a/core/storage/local/watcher_test.go b/core/storage/local/watcher_test.go index 8d2d31367..37387c76a 100644 --- a/core/storage/local/watcher_test.go +++ b/core/storage/local/watcher_test.go @@ -137,3 +137,18 @@ type noopExtractor struct{} func (s noopExtractor) Parse(files ...string) (map[string]metadata.Info, error) { return nil, nil } func (s noopExtractor) Version() string { return "0" } + +var _ = Describe("Watcher.Start", func() { + It("returns an error instead of hanging when the path cannot be watched", func() { + local.RegisterExtractor("noop", func(fs fs.FS, path string) local.Extractor { return noopExtractor{} }) + conf.Server.Scanner.Extractor = "noop" + ls, err := storage.For(filepath.Join(GinkgoT().TempDir(), "does-not-exist")) + Expect(err).ToNot(HaveOccurred()) + lsw := ls.(storage.Watcher) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + _, err = lsw.Start(ctx) + Expect(err).To(HaveOccurred()) + }) +}) diff --git a/core/stream/decider.go b/core/stream/decider.go index 3c6b01e05..38839cdf0 100644 --- a/core/stream/decider.go +++ b/core/stream/decider.go @@ -23,6 +23,7 @@ type TranscodeDecider interface { CreateTranscodeParams(decision *TranscodeDecision) (string, error) ResolveRequestFromToken(ctx context.Context, token string, mf *model.MediaFile, offset int) (Request, error) ResolveRequest(ctx context.Context, mf *model.MediaFile, reqFormat string, reqBitRate int, offset int) Request + ResolveClientRequest(ctx context.Context, mf *model.MediaFile, clientInfo *ClientInfo, offset int) Request } func NewTranscodeDecider(ds model.DataStore, ff ffmpeg.FFmpeg) TranscodeDecider { @@ -311,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 { @@ -326,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 } @@ -446,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/decider_test.go b/core/stream/decider_test.go index 577207636..01fef1249 100644 --- a/core/stream/decider_test.go +++ b/core/stream/decider_test.go @@ -1144,6 +1144,82 @@ var _ = Describe("Decider", func() { }) }) + Context("Player-forced format", func() { + symfonium := func() *ClientInfo { + return &ClientInfo{ + Name: "Symfonium", + DirectPlayProfiles: []DirectPlayProfile{ + {Containers: []string{"mp3", "flac", "ogg"}, Protocols: []string{ProtocolHTTP}}, + }, + TranscodingProfiles: []Profile{ + {Container: "flac", AudioCodec: "flac", Protocol: ProtocolHTTP}, + {Container: "mp3", AudioCodec: "mp3", Protocol: ProtocolHTTP}, + }, + } + } + + It("direct plays a flac source forced to flac", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "flac", Codec: "FLAC", BitRate: 1026, Channels: 2, SampleRate: 44100, BitDepth: new(16)}) + ci := symfonium() + Expect(ci.ForceFormat("flac")).To(BeTrue()) + + decision, err := svc.MakeDecision(ctx, mf, ci, TranscodeOptions{}) + Expect(err).ToNot(HaveOccurred()) + Expect(decision.CanDirectPlay).To(BeTrue()) + }) + + It("still transcodes a 24-bit flac when the client caps bit depth", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "flac", Codec: "FLAC", BitRate: 4600, Channels: 2, SampleRate: 96000, BitDepth: new(24)}) + ci := symfonium() + ci.CodecProfiles = []CodecProfile{{ + Type: CodecProfileTypeAudio, Name: "flac", + Limitations: []Limitation{{Name: LimitationAudioBitdepth, Comparison: ComparisonLessThanEqual, Values: []string{"16"}, Required: true}}, + }} + Expect(ci.ForceFormat("flac")).To(BeTrue()) + + decision, err := svc.MakeDecision(ctx, mf, ci, TranscodeOptions{}) + Expect(err).ToNot(HaveOccurred()) + Expect(decision.CanDirectPlay).To(BeFalse()) + Expect(decision.CanTranscode).To(BeTrue()) + Expect(decision.TranscodeStream.BitDepth).To(Equal(16)) + }) + + It("still transcodes a 320 mp3 forced to mp3 at a lower bitrate", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "mp3", Codec: "MP3", BitRate: 320, Channels: 2, SampleRate: 44100}) + ci := symfonium() + Expect(ci.ForceFormat("mp3")).To(BeTrue()) + ci.CapBitrate(192) + + decision, err := svc.MakeDecision(ctx, mf, ci, TranscodeOptions{}) + Expect(err).ToNot(HaveOccurred()) + Expect(decision.CanDirectPlay).To(BeFalse()) + Expect(decision.CanTranscode).To(BeTrue()) + Expect(decision.TargetBitrate).To(Equal(192)) + }) + + It("direct plays a 128 mp3 forced to mp3 at a higher bitrate", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "mp3", Codec: "MP3", BitRate: 128, Channels: 2, SampleRate: 44100}) + ci := symfonium() + Expect(ci.ForceFormat("mp3")).To(BeTrue()) + ci.CapBitrate(192) + + decision, err := svc.MakeDecision(ctx, mf, ci, TranscodeOptions{}) + Expect(err).ToNot(HaveOccurred()) + Expect(decision.CanDirectPlay).To(BeTrue()) + }) + + It("transcodes a flac source forced to mp3", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "flac", Codec: "FLAC", BitRate: 1026, Channels: 2, SampleRate: 44100, BitDepth: new(16)}) + ci := symfonium() + Expect(ci.ForceFormat("mp3")).To(BeTrue()) + + decision, err := svc.MakeDecision(ctx, mf, ci, TranscodeOptions{}) + Expect(err).ToNot(HaveOccurred()) + Expect(decision.CanDirectPlay).To(BeFalse()) + Expect(decision.CanTranscode).To(BeTrue()) + Expect(decision.TargetFormat).To(Equal("mp3")) + }) + }) }) Describe("ensureProbed", func() { diff --git a/core/stream/legacy_client.go b/core/stream/legacy_client.go index 652e42eba..5b732c3e2 100644 --- a/core/stream/legacy_client.go +++ b/core/stream/legacy_client.go @@ -73,37 +73,8 @@ func (s *deciderService) ResolveRequest(ctx context.Context, mf *model.MediaFile } clientInfo := buildLegacyClientInfo(mf, reqFormat, reqBitRate, playerMaxBitRate) - - // Apply server-side player transcoding override before making the decision - if trc, ok := request.TranscodingFrom(ctx); ok && trc.TargetFormat != "" { - clientInfo = applyServerOverride(ctx, clientInfo, &trc) - } else if player, ok := request.PlayerFrom(ctx); ok { - modified := *clientInfo - if modified.CapBitrate(player.MaxBitRate) { - clientInfo = &modified - log.Debug(ctx, "Applied player MaxBitRate cap", "playerMaxBitRate", player.MaxBitRate, "client", clientInfo.Name) - } - } - - decision, err := s.MakeDecision(ctx, mf, clientInfo, TranscodeOptions{SkipProbe: true}) - if err != nil { - log.Error(ctx, "Error making transcode decision, falling back to raw", "id", mf.ID, err) - req.Format = "raw" - return req - } - - if decision.CanDirectPlay { - req.Format = "raw" - return req - } - - if decision.CanTranscode { - req.Format = decision.TargetFormat - req.BitRate = decision.TargetBitrate - req.SampleRate = decision.TargetSampleRate - req.BitDepth = decision.TargetBitDepth - req.Channels = decision.TargetChannels - return req + if resolved, ok := s.resolve(ctx, mf, clientInfo, offset); ok { + return resolved } // No compatible profile for the requested format — retry with DefaultDownsamplingFormat @@ -119,3 +90,47 @@ func (s *deciderService) ResolveRequest(ctx context.Context, mf *model.MediaFile req.Format = "raw" return req } + +// ResolveClientRequest resolves a stream request for a client that declared its own direct play +// and transcoding profiles, falling back to raw when none fits. +func (s *deciderService) ResolveClientRequest(ctx context.Context, mf *model.MediaFile, clientInfo *ClientInfo, offset int) Request { + if req, ok := s.resolve(ctx, mf, clientInfo, offset); ok { + return req + } + return Request{Format: "raw", Offset: offset} +} + +// resolve applies the server-side player overrides to clientInfo and maps the decision to a +// Request. ok is false when no profile fits. +func (s *deciderService) resolve(ctx context.Context, mf *model.MediaFile, clientInfo *ClientInfo, offset int) (Request, bool) { + req := Request{Offset: offset} + if trc, ok := request.TranscodingFrom(ctx); ok && trc.TargetFormat != "" { + clientInfo = applyServerOverride(ctx, clientInfo, &trc) + } else if player, ok := request.PlayerFrom(ctx); ok { + modified := *clientInfo + if modified.CapBitrate(player.MaxBitRate) { + clientInfo = &modified + log.Debug(ctx, "Applied player MaxBitRate cap", "playerMaxBitRate", player.MaxBitRate, "client", clientInfo.Name) + } + } + + decision, err := s.MakeDecision(ctx, mf, clientInfo, TranscodeOptions{SkipProbe: true}) + if err != nil { + log.Error(ctx, "Error making transcode decision, falling back to raw", "id", mf.ID, err) + req.Format = "raw" + return req, true + } + switch { + case decision.CanDirectPlay: + req.Format = "raw" + case decision.CanTranscode: + req.Format = decision.TargetFormat + req.BitRate = decision.TargetBitrate + req.SampleRate = decision.TargetSampleRate + req.BitDepth = decision.TargetBitDepth + req.Channels = decision.TargetChannels + default: + return req, false + } + return req, true +} diff --git a/core/stream/legacy_client_test.go b/core/stream/legacy_client_test.go index bc8405976..ad3417faf 100644 --- a/core/stream/legacy_client_test.go +++ b/core/stream/legacy_client_test.go @@ -414,3 +414,72 @@ var _ = Describe("ResolveRequest", func() { }) }) }) + +var _ = Describe("ResolveClientRequest", func() { + var ( + svc TranscodeDecider + ctx context.Context + ) + mp3Target := []Profile{{Container: "mp3", AudioCodec: "mp3", Protocol: ProtocolHTTP}} + + BeforeEach(func() { + ctx = GinkgoT().Context() + ds := &tests.MockDataStore{ + MockedProperty: &tests.MockedPropertyRepo{}, + MockedTranscoding: &tests.MockTranscodingRepo{}, + } + auth.Init(ds) + svc = NewTranscodeDecider(ds, tests.NewMockFFmpeg("")) + }) + + It("direct plays a source matching a profile through container aliases", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "m4a", Codec: "AAC", BitRate: 256, Channels: 2, SampleRate: 44100}) + ci := &ClientInfo{ + DirectPlayProfiles: []DirectPlayProfile{{Containers: []string{"mp4"}, AudioCodecs: []string{"aac"}}}, + TranscodingProfiles: mp3Target, + } + + req := svc.ResolveClientRequest(ctx, mf, ci, 5) + + Expect(req.Format).To(Equal("raw")) + Expect(req.Offset).To(Equal(5)) + }) + + It("transcodes to the client's transcoding profile when no direct play profile matches", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "flac", Codec: "FLAC", BitRate: 1000, Channels: 2, SampleRate: 44100, BitDepth: new(16)}) + ci := &ClientInfo{ + DirectPlayProfiles: []DirectPlayProfile{{Containers: []string{"mp3"}}}, + TranscodingProfiles: mp3Target, + } + + Expect(svc.ResolveClientRequest(ctx, mf, ci, 0).Format).To(Equal("mp3")) + }) + + It("transcodes a direct-playable source over the bitrate cap to the client's target", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "flac", Codec: "FLAC", BitRate: 1000, Channels: 2, SampleRate: 44100, BitDepth: new(16)}) + ci := &ClientInfo{ + MaxAudioBitrate: 128, + DirectPlayProfiles: []DirectPlayProfile{{Containers: []string{"flac"}}}, + TranscodingProfiles: mp3Target, + } + + req := svc.ResolveClientRequest(ctx, mf, ci, 0) + + Expect(req.Format).To(Equal("mp3")) + Expect(req.BitRate).To(Equal(128)) + }) + + It("applies the player's MaxBitRate cap", func() { + mf := withProbe(&model.MediaFile{ID: "1", Suffix: "flac", Codec: "FLAC", BitRate: 1000, Channels: 2, SampleRate: 44100, BitDepth: new(16)}) + ctx = request.WithPlayer(ctx, model.Player{MaxBitRate: 192}) + ci := &ClientInfo{ + DirectPlayProfiles: []DirectPlayProfile{{Containers: []string{"flac"}}}, + TranscodingProfiles: mp3Target, + } + + req := svc.ResolveClientRequest(ctx, mf, ci, 0) + + Expect(req.Format).To(Equal("mp3")) + Expect(req.BitRate).To(Equal(192)) + }) +}) diff --git a/core/stream/media_streamer.go b/core/stream/media_streamer.go index aaa3126b4..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 @@ -268,6 +268,7 @@ func NewTranscodingCache() TranscodingCache { BitDepth: job.bitDepth, Channels: job.channels, Offset: job.offset, + Duration: job.mf.Duration, }) if err != nil { release() 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/stream/types.go b/core/stream/types.go index 19474dd91..300017c01 100644 --- a/core/stream/types.go +++ b/core/stream/types.go @@ -59,28 +59,38 @@ func (ci *ClientInfo) CapBitrate(maxKbps int) bool { return changed } -// ForceFormat narrows the client to transcoding to targetFormat and suppresses -// direct play, but only if the client already declares a profile for that -// format. All matching profiles are kept so negotiation can still pick among -// them (e.g. by protocol). Returns false (no-op) when targetFormat is empty or -// unsupported. +// ForceFormat narrows the client to transcoding to targetFormat, but only if the +// client already declares a profile for it. All matching profiles are kept so +// negotiation can still pick among them (e.g. by protocol). Direct play is rebuilt +// from those profiles rather than dropped, since declaring a transcoding profile +// for a format is proof the client can play it. Returns false when unsupported. func (ci *ClientInfo) ForceFormat(targetFormat string) bool { if targetFormat == "" { return false } var matched []Profile + var directPlay []DirectPlayProfile for i := range ci.TranscodingProfiles { + p := &ci.TranscodingProfiles[i] // matchesContainer is alias-aware, so a forced "oga" (legacy Opus // target_format) still matches a resolved "opus" profile. - if _, format := resolveTargetFormat(&ci.TranscodingProfiles[i]); matchesContainer(format, []string{targetFormat}) { - matched = append(matched, ci.TranscodingProfiles[i]) + container, format := resolveTargetFormat(p) + if !matchesContainer(format, []string{targetFormat}) { + continue } + matched = append(matched, *p) + directPlay = append(directPlay, DirectPlayProfile{ + Containers: []string{container}, + AudioCodecs: []string{format}, + Protocols: []string{ProtocolHTTP}, + MaxAudioChannels: p.MaxAudioChannels, + }) } if len(matched) == 0 { return false } ci.TranscodingProfiles = matched - ci.DirectPlayProfiles = nil + ci.DirectPlayProfiles = directPlay return true } diff --git a/core/stream/types_test.go b/core/stream/types_test.go index eff408362..88ad904a7 100644 --- a/core/stream/types_test.go +++ b/core/stream/types_test.go @@ -58,7 +58,7 @@ var _ = Describe("ClientInfo", func() { }) Describe("ForceFormat", func() { - It("restricts to the forced format and clears direct play when supported", func() { + It("restricts direct play to the forced format when supported", func() { ci := &ClientInfo{ DirectPlayProfiles: []DirectPlayProfile{{Containers: []string{"flac"}, AudioCodecs: []string{"flac"}}}, TranscodingProfiles: []Profile{ @@ -71,7 +71,35 @@ var _ = Describe("ClientInfo", func() { Expect(ok).To(BeTrue()) Expect(ci.TranscodingProfiles).To(HaveLen(1)) Expect(ci.TranscodingProfiles[0].AudioCodec).To(Equal("opus")) - Expect(ci.DirectPlayProfiles).To(BeEmpty()) + Expect(ci.DirectPlayProfiles).To(ConsistOf(DirectPlayProfile{ + Containers: []string{"ogg"}, AudioCodecs: []string{"opus"}, Protocols: []string{ProtocolHTTP}, + })) + }) + + It("keeps direct play for a source already in the forced format", func() { + ci := &ClientInfo{ + DirectPlayProfiles: []DirectPlayProfile{{Containers: []string{"flac"}, AudioCodecs: []string{"flac"}}}, + TranscodingProfiles: []Profile{ + {Container: "flac", AudioCodec: "flac", Protocol: ProtocolHTTP}, + {Container: "mp3", AudioCodec: "mp3", Protocol: ProtocolHTTP}, + }, + } + ok := ci.ForceFormat("flac") + Expect(ok).To(BeTrue()) + Expect(ci.DirectPlayProfiles).To(ConsistOf(DirectPlayProfile{ + Containers: []string{"flac"}, AudioCodecs: []string{"flac"}, Protocols: []string{ProtocolHTTP}, + })) + }) + + It("carries the channel limit of the forced profile into direct play", func() { + ci := &ClientInfo{ + TranscodingProfiles: []Profile{ + {Container: "flac", AudioCodec: "flac", Protocol: ProtocolHTTP, MaxAudioChannels: 2}, + }, + } + Expect(ci.ForceFormat("flac")).To(BeTrue()) + Expect(ci.DirectPlayProfiles).To(HaveLen(1)) + Expect(ci.DirectPlayProfiles[0].MaxAudioChannels).To(Equal(2)) }) It("matches a container-only forced format (mp3)", func() { 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/core/wire_providers.go b/core/wire_providers.go index a09fcc108..b3df9b2dc 100644 --- a/core/wire_providers.go +++ b/core/wire_providers.go @@ -10,6 +10,7 @@ import ( "github.com/navidrome/navidrome/core/metrics" "github.com/navidrome/navidrome/core/playback" "github.com/navidrome/navidrome/core/playlists" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/core/scrobbler" "github.com/navidrome/navidrome/core/stream" ) @@ -32,6 +33,7 @@ var Set = wire.NewSet( ffmpeg.New, scrobbler.GetPlayTracker, playback.GetInstance, + quickconnect.GetInstance, metrics.GetInstance, lyrics.NewLyrics, ) diff --git a/db/backup.go b/db/backup.go index 806bef8e2..74cae553e 100644 --- a/db/backup.go +++ b/db/backup.go @@ -9,6 +9,7 @@ import ( "path/filepath" "regexp" "slices" + "strings" "time" "github.com/mattn/go-sqlite3" @@ -18,7 +19,7 @@ import ( const ( backupPrefix = "navidrome_backup" - backupRegexString = backupPrefix + "_(.+)\\.db" + backupRegexString = "^" + backupPrefix + "_(.+)\\.db$" ) var backupRegex = regexp.MustCompile(backupRegexString) @@ -40,6 +41,18 @@ func backupOrRestore(ctx context.Context, isBackup bool, path string) error { } defer existingConn.Close() + // The driver opens with SQLITE_OPEN_CREATE, so without this check a typo in the + // path would create an empty database and "restore" it over the live one. + if !isBackup { + // The driver splits the DSN at '?', so such a path would open a different file. + if strings.ContainsRune(path, '?') { + return fmt.Errorf("backup path cannot contain '?': %s", path) + } + if _, err := os.Stat(path); err != nil { + return fmt.Errorf("backup file not available: %w", err) + } + } + backupDb, err := sql.Open(Driver, path) if err != nil { return fmt.Errorf("opening backup database in '%s': %w", path, err) diff --git a/db/backup_test.go b/db/backup_test.go index 5d1bfc6e3..609cba3d6 100644 --- a/db/backup_test.go +++ b/db/backup_test.go @@ -5,12 +5,14 @@ import ( "database/sql" "math/rand" "os" + "path/filepath" "time" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" . "github.com/navidrome/navidrome/db" "github.com/navidrome/navidrome/tests" + "github.com/navidrome/navidrome/utils/singleton" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -103,6 +105,19 @@ var _ = Describe("database backups", func() { Entry("delete all files", 0, 0), Entry("preserve all files when at length", len(timesDecreasingChronologically), len(timesDecreasingChronologically)), Entry("preserve all files when less than count", 10000, len(timesDecreasingChronologically))) + + It("ignores SQLite sidecar files when counting backups", func() { + for _, suffix := range []string{"-shm", "-wal"} { + file, err := os.Create(BackupPath(timesDecreasingChronologically[0]) + suffix) + Expect(err).ToNot(HaveOccurred()) + _ = file.Close() + } + + conf.Server.Backup.Count = len(timesDecreasingChronologically) + pruneCount, err := Prune(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(pruneCount).To(BeZero()) + }) }) Describe("backup and restore", Ordered, func() { @@ -148,4 +163,78 @@ var _ = Describe("database backups", func() { Expect(IsSchemaEmpty(ctx, Db())).To(BeFalse()) }) }) + + Describe("backup and restore with a file-based database", Ordered, func() { + var ctx context.Context + var tempFolder string + var dbFilePath string + + BeforeAll(func() { + ctx = context.Background() + DeferCleanup(configtest.SetupConfig()) + + var err error + tempFolder, err = os.MkdirTemp("", "navidrome_restore") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { + Close(ctx) + _ = os.RemoveAll(tempFolder) + }) + + // Mimic the production DSN (consts.DefaultDbPath): a file database in WAL mode. + dbFilePath = filepath.Join(tempFolder, "navidrome.db") + conf.Server.DbPath = dbFilePath + "?_busy_timeout=15000&_journal_mode=WAL&_foreign_keys=on&synchronous=normal" + // The previous container's cleanup closed the shared *sql.DB without + // dropping the singleton; force a fresh connection for this container. + singleton.DeleteInstance[*sql.DB]() + DeferCleanup(Init(ctx)) + }) + + It("restores data into a database whose stale WAL sidecar files were left behind", func() { + By("seeding user data in the current database") + _, err := Db().ExecContext(ctx, `INSERT INTO user (id, user_name, name, email, password, is_admin, created_at, updated_at) + VALUES ('u-restore-1', 'drilladmin', 'drilladmin', 'drilladmin@example.com', 'x', 1, datetime('now'), datetime('now'))`) + Expect(err).ToNot(HaveOccurred()) + + By("creating a backup containing the user row") + path, err := Backup(ctx) + Expect(err).ToNot(HaveOccurred()) + + By("simulating the CLI exiting without closing the pool: sidecar files stay behind") + _, err = Db().ExecContext(ctx, "CREATE TABLE IF NOT EXISTS _restore_probe(x)") + Expect(err).ToNot(HaveOccurred()) + singleton.DeleteInstance[*sql.DB]() + + err = tests.ClearDB() + Expect(err).ToNot(HaveOccurred()) + + By("restoring the backup") + Expect(Restore(ctx, path)).To(Succeed()) + + By("verifying the restored data is readable through a fresh connection") + singleton.DeleteInstance[*sql.DB]() + var userName string + Expect(Db().QueryRowContext(ctx, "SELECT user_name FROM user WHERE id = 'u-restore-1'").Scan(&userName)).To(Succeed()) + Expect(userName).To(Equal("drilladmin")) + }) + + It("fails to restore from a backup file that does not exist, leaving the database intact", func() { + By("seeding user data in the current database") + _, err := Db().ExecContext(ctx, `INSERT INTO user (id, user_name, name, email, password, is_admin, created_at, updated_at) + VALUES ('u-restore-2', 'keepme', 'keepme', 'keepme@example.com', 'x', 1, datetime('now'), datetime('now'))`) + Expect(err).ToNot(HaveOccurred()) + + By("attempting a restore from a nonexistent file") + missingPath := filepath.Join(tempFolder, "does_not_exist.db") + err = Restore(ctx, missingPath) + Expect(err).To(HaveOccurred()) + + By("verifying the database was not wiped") + var userName string + Expect(Db().QueryRowContext(ctx, "SELECT user_name FROM user WHERE id = 'u-restore-2'").Scan(&userName)).To(Succeed()) + Expect(userName).To(Equal("keepme")) + _, statErr := os.Stat(missingPath) + Expect(statErr).To(MatchError(os.ErrNotExist)) + }) + }) }) diff --git a/db/db.go b/db/db.go index c53aa364a..66e48cede 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...)) } @@ -157,17 +163,28 @@ func hasPendingMigrations(ctx context.Context, db *sql.DB, folder string) bool { return l.numPending > 0 } +// hasGooseTable reports whether goose's bookkeeping table exists, i.e. whether the +// database has ever been migrated. +func hasGooseTable(ctx context.Context, db *sql.DB) (bool, error) { + var name string + err := db.QueryRowContext(ctx, + "SELECT name FROM sqlite_master WHERE type='table' AND name='goose_db_version'").Scan(&name) + if errors.Is(err, sql.ErrNoRows) { + return false, nil + } + return err == nil, err +} + func isSchemaEmpty(ctx context.Context, db *sql.DB) bool { - rows, err := db.QueryContext(ctx, "SELECT name FROM sqlite_master WHERE type='table' AND name='goose_db_version';") // nolint:rowserrcheck + found, err := hasGooseTable(ctx, db) if err != nil { log.Fatal(ctx, "Database could not be opened!", err) } - defer rows.Close() - return !rows.Next() + return !found } type logAdapter struct { - ctx context.Context + ctx context.Context //nolint:containedctx // goose logger interface has no ctx silent bool } 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/export_test.go b/db/export_test.go index 02b88cd66..d2034bcdb 100644 --- a/db/export_test.go +++ b/db/export_test.go @@ -2,6 +2,10 @@ package db // Definitions for testing private methods var ( + EmbedMigrations = embedMigrations + FTSTables = ftsTables + FTSTriggerSuffixes = ftsTriggerSuffixes + FTSSearchMigration = ftsSearchMigration IsSchemaEmpty = isSchemaEmpty BackupPath = backupPath OptimizeDBAt = optimizeAt diff --git a/db/migrations/20240511220020_add_library_table.go b/db/migrations/20240511220020_add_library_table.go index 55b521ca9..d7b4aa5c2 100644 --- a/db/migrations/20240511220020_add_library_table.go +++ b/db/migrations/20240511220020_add_library_table.go @@ -3,7 +3,6 @@ package migrations import ( "context" "database/sql" - "fmt" "github.com/navidrome/navidrome/conf" "github.com/pressly/goose/v3" @@ -28,10 +27,10 @@ func upAddLibraryTable(ctx context.Context, tx *sql.Tx) error { return err } - _, err = tx.ExecContext(ctx, fmt.Sprintf(` - insert into library(id, name, path) values(1, 'Music Library', '%s'); - delete from property where id like 'LastScan-%%'; -`, conf.Server.MusicFolder)) + _, err = tx.ExecContext(ctx, ` + insert into library(id, name, path) values(1, 'Music Library', ?); + delete from property where id like 'LastScan-%'; +`, conf.Server.MusicFolder) if err != nil { return err } 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/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/db/migrations/20260929221042_add_library_pid_columns.sql b/db/migrations/20260929221042_add_library_pid_columns.sql new file mode 100644 index 000000000..487512287 --- /dev/null +++ b/db/migrations/20260929221042_add_library_pid_columns.sql @@ -0,0 +1,18 @@ +-- +goose Up +-- +goose StatementBegin +alter table library add column pid_album varchar default '' not null; +alter table library add column pid_track varchar default '' not null; +alter table library add column scanned_pid_album varchar default '' not null; +alter table library add column scanned_pid_track varchar default '' not null; + +-- Every library was scanned with the global PID config, so seed it as their scanned config. +-- This way the upgrade does not trigger a full rescan. +update library set + scanned_pid_album = coalesce((select value from property where id = 'PIDAlbum'), ''), + scanned_pid_track = coalesce((select value from property where id = 'PIDTrack'), ''); + +delete from property where id in ('PIDAlbum', 'PIDTrack'); +-- +goose StatementEnd + +-- +goose Down +SELECT 1; 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/db/repair.go b/db/repair.go new file mode 100644 index 000000000..7f8b8d044 --- /dev/null +++ b/db/repair.go @@ -0,0 +1,367 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + "strings" +) + +var ftsTables = []string{"media_file_fts", "album_fts", "artist_fts"} + +var ftsTriggerSuffixes = []string{"_ai", "_ad", "_au"} + +// integrityCheckMaxIssues bounds the problems reported; IntegrityCheck asks the +// pragma for one extra row, because it truncates without emitting any marker. +const integrityCheckMaxIssues = 100 + +// IntegrityCheck runs PRAGMA integrity_check and returns the problems it reports, or +// an empty slice when healthy. The second value marks a list that was cut short. +func IntegrityCheck(ctx context.Context, database *sql.DB) ([]string, bool, error) { + rows, err := database.QueryContext(ctx, + fmt.Sprintf("PRAGMA integrity_check(%d)", integrityCheckMaxIssues+1)) + if err != nil { + return nil, false, fmt.Errorf("running integrity_check: %w", err) + } + defer rows.Close() + + var issues []string + for rows.Next() { + var line string + if err := rows.Scan(&line); err != nil { + return nil, false, fmt.Errorf("reading integrity_check results: %w", err) + } + issues = append(issues, line) + } + if err := rows.Err(); err != nil { + return nil, false, fmt.Errorf("reading integrity_check results: %w", err) + } + if len(issues) == 1 && issues[0] == "ok" { + return nil, false, nil + } + if len(issues) > integrityCheckMaxIssues { + return issues[:integrityCheckMaxIssues], true, nil + } + return issues, false, nil +} + +// FKViolation counts the rows in Table that reference missing rows in Parent. +type FKViolation struct { + Table string + Parent string + Count int64 +} + +// ForeignKeyCheck runs PRAGMA foreign_key_check, aggregated per (table, parent) pair +// because the raw pragma emits one row per orphan, unbounded on a large library. +func ForeignKeyCheck(ctx context.Context, database *sql.DB) ([]FKViolation, error) { + rows, err := database.QueryContext(ctx, + `SELECT "table", "parent", count(*) FROM pragma_foreign_key_check GROUP BY "table", "parent"`) + if err != nil { + return nil, fmt.Errorf("running foreign_key_check: %w", err) + } + defer rows.Close() + + var violations []FKViolation + for rows.Next() { + var v FKViolation + if err := rows.Scan(&v.Table, &v.Parent, &v.Count); err != nil { + return nil, fmt.Errorf("reading foreign_key_check results: %w", err) + } + violations = append(violations, v) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("reading foreign_key_check results: %w", err) + } + return violations, nil +} + +// IsFTSCorruptionOnly reports whether every integrity issue refers to one of the +// FTS5 search tables, meaning RebuildFTS can fully repair the database. +func IsFTSCorruptionOnly(issues []string) bool { + if len(issues) == 0 { + return false + } + for _, line := range issues { + if !slices.ContainsFunc(ftsTables, func(table string) bool { return strings.Contains(line, table) }) { + return false + } + } + return true +} + +// execer is the subset of *sql.DB and *sql.Tx that verifyFTS needs. +type execer interface { + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) +} + +// VerifyFTS runs the FTS5 'integrity-check' command on each search table. Unlike a +// full PRAGMA integrity_check, it reads only the FTS indexes, not the whole database. +func VerifyFTS(ctx context.Context, database *sql.DB) error { + return verifyFTS(ctx, database) +} + +func verifyFTS(ctx context.Context, database execer) error { + for _, table := range ftsTables { + stmt := fmt.Sprintf("INSERT INTO %[1]s(%[1]s) VALUES('integrity-check')", table) //nolint:gosec // fixed table list + if _, err := database.ExecContext(ctx, stmt); err != nil { + return fmt.Errorf("verifying %s: %w", table, err) + } + } + return nil +} + +const ftsSearchMigration int64 = 20260220173400 + +var errNotMigrated = errors.New("the FTS search migration has not been applied yet; start Navidrome once to migrate the database first") + +// requireFTSMigration fails unless the FTS search migration has run. The goose table +// is probed separately because a query against a missing table fails at prepare time. +func requireFTSMigration(ctx context.Context, database *sql.DB) error { + migrated, err := hasGooseTable(ctx, database) + if err != nil { + return fmt.Errorf("checking FTS migration status: %w", err) + } + if !migrated { + return errNotMigrated + } + var applied int + if err := database.QueryRowContext(ctx, + "SELECT count(*) FROM goose_db_version WHERE version_id = ?", ftsSearchMigration).Scan(&applied); err != nil { + return fmt.Errorf("checking FTS migration status: %w", err) + } + if applied == 0 { + return errNotMigrated + } + return nil +} + +// RebuildFTS drops the FTS5 search tables and their triggers, recreates them from the +// base tables, and verifies the result before committing. The tables are contentless, +// so no user data is lost. It needs only the FTS migration, not a fully migrated +// schema, because a corrupted DB often cannot run pending migrations. +func RebuildFTS(ctx context.Context, database *sql.DB) error { + if err := requireFTSMigration(ctx, database); err != nil { + return err + } + + tx, err := database.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("starting FTS rebuild transaction: %w", err) + } + defer func() { _ = tx.Rollback() }() + + var stmts []string + for _, table := range ftsTables { + for _, suffix := range ftsTriggerSuffixes { + stmts = append(stmts, "DROP TRIGGER IF EXISTS "+table+suffix) + } + stmts = append(stmts, "DROP TABLE IF EXISTS "+table) + } + stmts = append(stmts, ftsSchemaDDL...) + for _, stmt := range stmts { + if _, err := tx.ExecContext(ctx, stmt); err != nil { + return fmt.Errorf("rebuilding FTS schema: %w", err) + } + } + if err := verifyFTS(ctx, tx); err != nil { + return fmt.Errorf("the rebuilt search index did not verify: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("committing FTS rebuild: %w", err) + } + return nil +} + +// ftsSchemaDDL must reproduce what the full migration chain produces, not what any +// single migration does; the schema comparison in repair_test.go guards the drift. +var ftsSchemaDDL = []string{ + ` + CREATE VIRTUAL TABLE IF NOT EXISTS media_file_fts USING fts5( + title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + disc_subtitle, search_participants, search_normalized, + content='', content_rowid='rowid', + tokenize='unicode61 remove_diacritics 2' + ) + `, + ` + CREATE VIRTUAL TABLE IF NOT EXISTS album_fts USING fts5( + name, sort_album_name, album_artist, + search_participants, discs, catalog_num, album_version, search_normalized, + content='', content_rowid='rowid', + tokenize='unicode61 remove_diacritics 2' + ) + `, + ` + CREATE VIRTUAL TABLE IF NOT EXISTS artist_fts USING fts5( + name, sort_artist_name, search_normalized, + content='', content_rowid='rowid', + tokenize='unicode61 remove_diacritics 2' + ) + `, + ` + INSERT INTO media_file_fts(rowid, title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + disc_subtitle, search_participants, search_normalized) + SELECT rowid, title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + COALESCE(disc_subtitle, ''), COALESCE(search_participants, ''), + COALESCE(search_normalized, '') + FROM media_file + `, + ` + INSERT INTO album_fts(rowid, name, sort_album_name, album_artist, + search_participants, discs, catalog_num, album_version, search_normalized) + SELECT rowid, name, COALESCE(sort_album_name, ''), COALESCE(album_artist, ''), + COALESCE(search_participants, ''), COALESCE(discs, ''), + COALESCE(catalog_num, ''), + COALESCE((SELECT group_concat(json_extract(je.value, '$.value'), ' ') + FROM json_each(album.tags, '$.albumversion') AS je), ''), + COALESCE(search_normalized, '') + FROM album + `, + ` + INSERT INTO artist_fts(rowid, name, sort_artist_name, search_normalized) + SELECT rowid, name, COALESCE(sort_artist_name, ''), COALESCE(search_normalized, '') + FROM artist + `, + ` + CREATE TRIGGER media_file_fts_ai AFTER INSERT ON media_file BEGIN + INSERT INTO media_file_fts(rowid, title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + disc_subtitle, search_participants, search_normalized) + VALUES (NEW.rowid, NEW.title, NEW.album, NEW.artist, NEW.album_artist, + NEW.sort_title, NEW.sort_album_name, NEW.sort_artist_name, NEW.sort_album_artist_name, + COALESCE(NEW.disc_subtitle, ''), COALESCE(NEW.search_participants, ''), + COALESCE(NEW.search_normalized, '')); + END + `, + ` + CREATE TRIGGER media_file_fts_ad AFTER DELETE ON media_file BEGIN + INSERT INTO media_file_fts(media_file_fts, rowid, title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + disc_subtitle, search_participants, search_normalized) + VALUES ('delete', OLD.rowid, OLD.title, OLD.album, OLD.artist, OLD.album_artist, + OLD.sort_title, OLD.sort_album_name, OLD.sort_artist_name, OLD.sort_album_artist_name, + COALESCE(OLD.disc_subtitle, ''), COALESCE(OLD.search_participants, ''), + COALESCE(OLD.search_normalized, '')); + END + `, + ` + CREATE TRIGGER media_file_fts_au AFTER UPDATE ON media_file + WHEN + OLD.title IS NOT NEW.title OR + OLD.album IS NOT NEW.album OR + OLD.artist IS NOT NEW.artist OR + OLD.album_artist IS NOT NEW.album_artist OR + OLD.sort_title IS NOT NEW.sort_title OR + OLD.sort_album_name IS NOT NEW.sort_album_name OR + OLD.sort_artist_name IS NOT NEW.sort_artist_name OR + OLD.sort_album_artist_name IS NOT NEW.sort_album_artist_name OR + OLD.disc_subtitle IS NOT NEW.disc_subtitle OR + OLD.search_participants IS NOT NEW.search_participants OR + OLD.search_normalized IS NOT NEW.search_normalized + BEGIN + INSERT INTO media_file_fts(media_file_fts, rowid, title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + disc_subtitle, search_participants, search_normalized) + VALUES ('delete', OLD.rowid, OLD.title, OLD.album, OLD.artist, OLD.album_artist, + OLD.sort_title, OLD.sort_album_name, OLD.sort_artist_name, OLD.sort_album_artist_name, + COALESCE(OLD.disc_subtitle, ''), COALESCE(OLD.search_participants, ''), + COALESCE(OLD.search_normalized, '')); + INSERT INTO media_file_fts(rowid, title, album, artist, album_artist, + sort_title, sort_album_name, sort_artist_name, sort_album_artist_name, + disc_subtitle, search_participants, search_normalized) + VALUES (NEW.rowid, NEW.title, NEW.album, NEW.artist, NEW.album_artist, + NEW.sort_title, NEW.sort_album_name, NEW.sort_artist_name, NEW.sort_album_artist_name, + COALESCE(NEW.disc_subtitle, ''), COALESCE(NEW.search_participants, ''), + COALESCE(NEW.search_normalized, '')); + END + `, + ` + CREATE TRIGGER album_fts_ai AFTER INSERT ON album BEGIN + INSERT INTO album_fts(rowid, name, sort_album_name, album_artist, + search_participants, discs, catalog_num, album_version, search_normalized) + VALUES (NEW.rowid, NEW.name, COALESCE(NEW.sort_album_name, ''), COALESCE(NEW.album_artist, ''), + COALESCE(NEW.search_participants, ''), COALESCE(NEW.discs, ''), + COALESCE(NEW.catalog_num, ''), + COALESCE((SELECT group_concat(json_extract(je.value, '$.value'), ' ') + FROM json_each(NEW.tags, '$.albumversion') AS je), ''), + COALESCE(NEW.search_normalized, '')); + END + `, + ` + CREATE TRIGGER album_fts_ad AFTER DELETE ON album BEGIN + INSERT INTO album_fts(album_fts, rowid, name, sort_album_name, album_artist, + search_participants, discs, catalog_num, album_version, search_normalized) + VALUES ('delete', OLD.rowid, OLD.name, COALESCE(OLD.sort_album_name, ''), COALESCE(OLD.album_artist, ''), + COALESCE(OLD.search_participants, ''), COALESCE(OLD.discs, ''), + COALESCE(OLD.catalog_num, ''), + COALESCE((SELECT group_concat(json_extract(je.value, '$.value'), ' ') + FROM json_each(OLD.tags, '$.albumversion') AS je), ''), + COALESCE(OLD.search_normalized, '')); + END + `, + ` + CREATE TRIGGER album_fts_au AFTER UPDATE ON album + WHEN + OLD.name IS NOT NEW.name OR + OLD.sort_album_name IS NOT NEW.sort_album_name OR + OLD.album_artist IS NOT NEW.album_artist OR + OLD.search_participants IS NOT NEW.search_participants OR + OLD.discs IS NOT NEW.discs OR + OLD.catalog_num IS NOT NEW.catalog_num OR + OLD.tags IS NOT NEW.tags OR + OLD.search_normalized IS NOT NEW.search_normalized + BEGIN + INSERT INTO album_fts(album_fts, rowid, name, sort_album_name, album_artist, + search_participants, discs, catalog_num, album_version, search_normalized) + VALUES ('delete', OLD.rowid, OLD.name, COALESCE(OLD.sort_album_name, ''), COALESCE(OLD.album_artist, ''), + COALESCE(OLD.search_participants, ''), COALESCE(OLD.discs, ''), + COALESCE(OLD.catalog_num, ''), + COALESCE((SELECT group_concat(json_extract(je.value, '$.value'), ' ') + FROM json_each(OLD.tags, '$.albumversion') AS je), ''), + COALESCE(OLD.search_normalized, '')); + INSERT INTO album_fts(rowid, name, sort_album_name, album_artist, + search_participants, discs, catalog_num, album_version, search_normalized) + VALUES (NEW.rowid, NEW.name, COALESCE(NEW.sort_album_name, ''), COALESCE(NEW.album_artist, ''), + COALESCE(NEW.search_participants, ''), COALESCE(NEW.discs, ''), + COALESCE(NEW.catalog_num, ''), + COALESCE((SELECT group_concat(json_extract(je.value, '$.value'), ' ') + FROM json_each(NEW.tags, '$.albumversion') AS je), ''), + COALESCE(NEW.search_normalized, '')); + END + `, + ` + CREATE TRIGGER artist_fts_ai AFTER INSERT ON artist BEGIN + INSERT INTO artist_fts(rowid, name, sort_artist_name, search_normalized) + VALUES (NEW.rowid, NEW.name, COALESCE(NEW.sort_artist_name, ''), + COALESCE(NEW.search_normalized, '')); + END + `, + ` + CREATE TRIGGER artist_fts_ad AFTER DELETE ON artist BEGIN + INSERT INTO artist_fts(artist_fts, rowid, name, sort_artist_name, search_normalized) + VALUES ('delete', OLD.rowid, OLD.name, COALESCE(OLD.sort_artist_name, ''), + COALESCE(OLD.search_normalized, '')); + END + `, + ` + CREATE TRIGGER artist_fts_au AFTER UPDATE ON artist + WHEN + OLD.name IS NOT NEW.name OR + OLD.sort_artist_name IS NOT NEW.sort_artist_name OR + OLD.search_normalized IS NOT NEW.search_normalized + BEGIN + INSERT INTO artist_fts(artist_fts, rowid, name, sort_artist_name, search_normalized) + VALUES ('delete', OLD.rowid, OLD.name, COALESCE(OLD.sort_artist_name, ''), + COALESCE(OLD.search_normalized, '')); + INSERT INTO artist_fts(rowid, name, sort_artist_name, search_normalized) + VALUES (NEW.rowid, NEW.name, COALESCE(NEW.sort_artist_name, ''), + COALESCE(NEW.search_normalized, '')); + END + `, +} diff --git a/db/repair_test.go b/db/repair_test.go new file mode 100644 index 000000000..eb7321939 --- /dev/null +++ b/db/repair_test.go @@ -0,0 +1,309 @@ +package db_test + +import ( + "context" + "database/sql" + "fmt" + "path/filepath" + "regexp" + "strings" + + "github.com/navidrome/navidrome/db" + "github.com/pressly/goose/v3" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// newDB returns an in-memory database migrated up to the given goose version +// (0 = fully migrated). +func newDB(ctx context.Context, upTo int64) *sql.DB { + GinkgoHelper() + d, err := sql.Open(db.Dialect, "file::memory:") + Expect(err).ToNot(HaveOccurred()) + d.SetMaxOpenConns(1) // non-shared :memory:, a second conn would be an empty DB + DeferCleanup(func() { _ = d.Close() }) + + _, err = d.ExecContext(ctx, "PRAGMA foreign_keys=off") + Expect(err).ToNot(HaveOccurred()) + goose.SetBaseFS(db.EmbedMigrations) + goose.SetLogger(goose.NopLogger()) + DeferCleanup(func() { goose.SetBaseFS(nil) }) + Expect(goose.SetDialect(db.Dialect)).To(Succeed()) + if upTo == 0 { + Expect(goose.UpContext(ctx, d, "migrations")).To(Succeed()) + } else { + Expect(goose.UpToContext(ctx, d, "migrations", upTo)).To(Succeed()) + } + return d +} + +// openMismatchedIndexDB builds a database whose index is declared over a different +// column than the one it was populated from, so integrity_check reports one issue per row. +func openMismatchedIndexDB(ctx context.Context, rows int) *sql.DB { + GinkgoHelper() + path := filepath.Join(GinkgoT().TempDir(), "mismatched.db") + open := func() *sql.DB { + d, err := sql.Open(db.Dialect, path) + Expect(err).ToNot(HaveOccurred()) + d.SetMaxOpenConns(1) + return d + } + d := open() + for _, stmt := range []string{ + `create table t(a, b)`, + fmt.Sprintf(`with recursive s(x) as (select 1 union all select x+1 from s where x < %d) + insert into t select x, x + 10000 from s`, rows), + `create index i on t(a)`, + `pragma writable_schema=on`, + `update sqlite_master set sql = 'CREATE INDEX i ON t(b)' where name = 'i'`, + } { + _, err := d.ExecContext(ctx, stmt) + Expect(err).ToNot(HaveOccurred()) + } + Expect(d.Close()).To(Succeed()) // reopen so SQLite reparses the doctored schema + + d = open() + DeferCleanup(func() { _ = d.Close() }) + return d +} + +var _ = Describe("IsFTSCorruptionOnly", func() { + It("is true when every issue mentions an FTS search table", func() { + Expect(db.IsFTSCorruptionOnly([]string{ + `fts5: corruption found reading blob 42 from table "media_file_fts"`, + `malformed inverted index for FTS5 table main.album_fts`, + `fts5: corruption in "artist_fts"`, + })).To(BeTrue()) + }) + + It("is false when any issue is outside the FTS search tables", func() { + Expect(db.IsFTSCorruptionOnly([]string{ + `fts5: corruption found reading blob 42 from table "media_file_fts"`, + `*** in database main ***`, + })).To(BeFalse()) + }) + + It("is false when there are no issues", func() { + Expect(db.IsFTSCorruptionOnly(nil)).To(BeFalse()) + }) +}) + +var _ = Describe("RebuildFTS schema guard", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = context.Background() + }) + + It("refuses to run on a schema older than the FTS migration", func() { + old := newDB(ctx, db.FTSSearchMigration-1) + + err := db.RebuildFTS(ctx, old) + Expect(err).To(MatchError(ContainSubstring("migration"))) + }) + + It("refuses to run on a database that was never migrated", func() { + empty, err := sql.Open(db.Dialect, "file::memory:") + Expect(err).ToNot(HaveOccurred()) + empty.SetMaxOpenConns(1) + DeferCleanup(func() { _ = empty.Close() }) + + Expect(db.RebuildFTS(ctx, empty)).To(MatchError(ContainSubstring("start Navidrome once"))) + }) + + It("runs on a post-FTS schema even when newer migrations are pending", func() { + behind := newDB(ctx, 20260702152457) + + Expect(db.RebuildFTS(ctx, behind)).To(Succeed()) + }) +}) + +var _ = Describe("Repair", func() { + var ( + ctx context.Context + database *sql.DB + ) + + BeforeEach(func() { + ctx = context.Background() + database = newDB(ctx, 0) + + for _, stmt := range []string{ + `insert into artist(id, name, search_normalized) values ('ar-1', 'Ramones', 'ramones')`, + `insert into album(id, name, search_normalized) values ('al-1', 'Rocket to Russia', 'rocket to russia')`, + `insert into media_file(id, title, search_normalized) values ('mf-1', 'Teenage Lobotomy', 'teenage lobotomy')`, + `insert into media_file(id, title, search_normalized) values ('mf-2', 'Rockaway Beach', 'rockaway beach')`, + } { + _, err := database.ExecContext(ctx, stmt) + Expect(err).ToNot(HaveOccurred()) + } + }) + + corruptFTS := func(table string) { + // 8+ bytes of garbage: a 4-byte blob still parses as a valid empty structure record + _, err := database.ExecContext(ctx, `update `+table+`_data set block = x'deadbeefdeadbeef' where id > 1`) //nolint:gosec + Expect(err).ToNot(HaveOccurred()) + } + + searchFTS := func(table, term string) int { + var count int + err := database.QueryRowContext(ctx, `select count(*) from `+table+` where `+table+` match ?`, term).Scan(&count) + Expect(err).ToNot(HaveOccurred()) + return count + } + + // ftsSchema returns name -> whitespace-normalized DDL for the FTS tables, + // their shadow tables, and their triggers. + ftsSchema := func() map[string]string { + rows, err := database.QueryContext(ctx, + `select name, sql from sqlite_master where name like '%_fts%' and sql is not null`) + Expect(err).ToNot(HaveOccurred()) + defer rows.Close() + ws := regexp.MustCompile(`\s+`) + schema := map[string]string{} + for rows.Next() { + var name, ddl string + Expect(rows.Scan(&name, &ddl)).To(Succeed()) + schema[name] = ws.ReplaceAllString(ddl, " ") + } + Expect(rows.Err()).ToNot(HaveOccurred()) + return schema + } + + Describe("IntegrityCheck", func() { + It("returns no issues for a healthy database", func() { + issues, truncated, err := db.IntegrityCheck(ctx, database) + Expect(err).ToNot(HaveOccurred()) + Expect(issues).To(BeEmpty()) + Expect(truncated).To(BeFalse()) + }) + + It("reports corruption in an FTS index", func() { + corruptFTS("media_file_fts") + issues, truncated, err := db.IntegrityCheck(ctx, database) + Expect(err).ToNot(HaveOccurred()) + Expect(issues).ToNot(BeEmpty()) + Expect(strings.Join(issues, "\n")).To(ContainSubstring("media_file_fts")) + Expect(truncated).To(BeFalse()) + }) + + It("flags the issue list as truncated when there are more issues than the limit", func() { + broken := openMismatchedIndexDB(ctx, 300) + + issues, truncated, err := db.IntegrityCheck(ctx, broken) + Expect(err).ToNot(HaveOccurred()) + Expect(issues).To(HaveLen(100)) + Expect(truncated).To(BeTrue()) + }) + + It("does not flag truncation when the issues exactly fill the limit", func() { + broken := openMismatchedIndexDB(ctx, 100) + + issues, truncated, err := db.IntegrityCheck(ctx, broken) + Expect(err).ToNot(HaveOccurred()) + Expect(issues).To(HaveLen(100)) + Expect(truncated).To(BeFalse()) + }) + }) + + Describe("ForeignKeyCheck", func() { + It("returns no violations for a healthy database", func() { + violations, err := db.ForeignKeyCheck(ctx, database) + Expect(err).ToNot(HaveOccurred()) + Expect(violations).To(BeEmpty()) + }) + + It("reports rows referencing missing parents", func() { + _, err := database.ExecContext(ctx, + `insert into media_file(id, title, library_id) values ('mf-bad', 'Orphan', 999)`) + Expect(err).ToNot(HaveOccurred()) + + violations, err := db.ForeignKeyCheck(ctx, database) + Expect(err).ToNot(HaveOccurred()) + Expect(violations).To(HaveLen(1)) + Expect(violations[0].Table).To(Equal("media_file")) + Expect(violations[0].Parent).To(Equal("library")) + Expect(violations[0].Count).To(BeNumerically("==", 1)) + }) + }) + + Describe("VerifyFTS", func() { + It("passes on a healthy index", func() { + Expect(db.VerifyFTS(ctx, database)).To(Succeed()) + }) + + It("fails on a corrupted index, naming the table", func() { + corruptFTS("album_fts") + Expect(db.VerifyFTS(ctx, database)).To(MatchError(ContainSubstring("album_fts"))) + }) + }) + + Describe("RebuildFTS", func() { + It("repairs a corrupted FTS index", func() { + corruptFTS("media_file_fts") + + Expect(db.RebuildFTS(ctx, database)).To(Succeed()) + + issues, _, err := db.IntegrityCheck(ctx, database) + Expect(err).ToNot(HaveOccurred()) + Expect(issues).To(BeEmpty()) + Expect(db.VerifyFTS(ctx, database)).To(Succeed()) + Expect(searchFTS("media_file_fts", "lobotomy")).To(Equal(1)) + Expect(searchFTS("album_fts", "russia")).To(Equal(1)) + Expect(searchFTS("artist_fts", "ramones")).To(Equal(1)) + }) + + It("recreates tables and triggers dropped by hand", func() { + for _, table := range db.FTSTables { + for _, suffix := range db.FTSTriggerSuffixes { + _, err := database.ExecContext(ctx, "drop trigger "+table+suffix) + Expect(err).ToNot(HaveOccurred()) + } + _, err := database.ExecContext(ctx, "drop table "+table) + Expect(err).ToNot(HaveOccurred()) + } + + Expect(db.RebuildFTS(ctx, database)).To(Succeed()) + + Expect(searchFTS("media_file_fts", "rockaway")).To(Equal(1)) + }) + + It("rolls back and keeps the old index when the rebuild fails", func() { + // Triggers go first: SQLite refuses to drop a column they reference. + for _, suffix := range db.FTSTriggerSuffixes { + _, err := database.ExecContext(ctx, "drop trigger media_file_fts"+suffix) + Expect(err).ToNot(HaveOccurred()) + } + _, err := database.ExecContext(ctx, `alter table media_file drop column disc_subtitle`) + Expect(err).ToNot(HaveOccurred()) + + Expect(db.RebuildFTS(ctx, database)).ToNot(Succeed()) + + Expect(searchFTS("media_file_fts", "lobotomy")).To(Equal(1)) + Expect(searchFTS("album_fts", "russia")).To(Equal(1)) + }) + + It("produces the same schema as the migration", func() { + migrated := ftsSchema() + Expect(migrated).ToNot(BeEmpty()) + + Expect(db.RebuildFTS(ctx, database)).To(Succeed()) + + Expect(ftsSchema()).To(Equal(migrated)) + }) + + It("leaves working triggers behind", func() { + Expect(db.RebuildFTS(ctx, database)).To(Succeed()) + + _, err := database.ExecContext(ctx, + `insert into artist(id, name, search_normalized) values ('ar-2', 'Blondie', 'blondie')`) + Expect(err).ToNot(HaveOccurred()) + Expect(searchFTS("artist_fts", "blondie")).To(Equal(1)) + + _, err = database.ExecContext(ctx, `delete from artist where id = 'ar-2'`) + Expect(err).ToNot(HaveOccurred()) + Expect(searchFTS("artist_fts", "blondie")).To(BeZero()) + }) + }) +}) diff --git a/go.mod b/go.mod index 4339b9c55..e96b8c8b3 100644 --- a/go.mod +++ b/go.mod @@ -3,30 +3,31 @@ module github.com/navidrome/navidrome go 1.27 // Fork to implement raw tags support -replace go.senan.xyz/taglib => github.com/deluan/go-taglib v0.0.0-20260905051825-df1d035571df +replace go.senan.xyz/taglib => github.com/deluan/go-taglib v0.0.0-20260913142955-d55e0c9353cb require ( github.com/Masterminds/squirrel v1.5.4 - github.com/andybalholm/cascadia v1.3.4 - github.com/bmatcuk/doublestar/v4 v4.10.0 - github.com/deluan/rest v0.0.0-20211102003136-6260bc399cbf + github.com/andybalholm/cascadia v1.3.5 + 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 github.com/djherbis/atime v1.1.0 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 + 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 github.com/go-chi/jwtauth/v5 v5.4.0 github.com/go-viper/encoding/ini v0.1.1 github.com/go-viper/mapstructure/v2 v2.5.0 - github.com/gohugoio/hashstructure v1.0.0 + github.com/gohugoio/hashstructure v1.1.0 github.com/google/go-pipeline v0.0.0-20230411140531-6cbedfc1d3fc github.com/google/uuid v1.6.0 github.com/google/wire v0.7.0 @@ -35,16 +36,16 @@ require ( github.com/jellydator/ttlcache/v3 v3.4.1 github.com/kardianos/service v1.3.0 github.com/kr/pretty v0.3.1 - github.com/lestrrat-go/jwx/v3 v3.2.0 - github.com/mattn/go-sqlite3 v1.14.50 + github.com/lestrrat-go/jwx/v3 v3.3.0 + 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.1 - 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 - github.com/pressly/goose/v3 v3.27.3 + github.com/pressly/goose/v3 v3.28.0 github.com/prometheus/client_golang v1.24.1 github.com/rjeczalik/notify v0.9.3 github.com/robfig/cron/v3 v3.0.1 @@ -60,13 +61,13 @@ require ( github.com/zeebo/xxh3 v1.1.0 go.senan.xyz/taglib v0.11.1 go.uber.org/goleak v1.3.0 - golang.org/x/image v0.45.0 - golang.org/x/net v0.58.0 - golang.org/x/sync v0.22.0 - golang.org/x/sys v0.47.0 - golang.org/x/term v0.45.0 - golang.org/x/text v0.41.0 - golang.org/x/time v0.15.0 + golang.org/x/image v0.46.0 + golang.org/x/net v0.59.0 + golang.org/x/sync v0.23.0 + golang.org/x/sys v0.48.0 + golang.org/x/term v0.46.0 + golang.org/x/text v0.42.0 + golang.org/x/time v0.16.0 gopkg.in/yaml.v3 v3.0.1 ) @@ -81,17 +82,20 @@ 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-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 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/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,11 +114,13 @@ 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 github.com/prometheus/common v0.70.1 // indirect - github.com/prometheus/procfs v0.21.1 // indirect + github.com/prometheus/procfs v0.22.0 // indirect github.com/rogpeppe/go-internal v1.16.0 // indirect github.com/sagikazarmark/locafero v0.12.0 // indirect github.com/sanity-io/litter v1.5.8 // indirect @@ -131,10 +137,10 @@ require ( go.opentelemetry.io/proto/otlp v1.11.0 // indirect go.uber.org/multierr v1.11.0 // indirect go.yaml.in/yaml/v3 v3.0.5 // indirect - golang.org/x/crypto v0.55.0 // indirect - golang.org/x/mod v0.40.0 // indirect - golang.org/x/telemetry v0.0.0-20260811182544-a038080d80e5 // indirect - golang.org/x/tools v0.49.0 // indirect + golang.org/x/crypto v0.57.0 // indirect + golang.org/x/mod v0.41.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 71d9facfd..fe6dbbc3f 100644 --- a/go.sum +++ b/go.sum @@ -6,16 +6,16 @@ github.com/Masterminds/semver/v3 v3.5.0 h1:kQceYJfbupGfZOKZQg0kou0DgAKhzDg2NZPAw github.com/Masterminds/semver/v3 v3.5.0/go.mod h1:4V+yj/TJE1HU9XfppCwVMZq3I84lprf4nC11bSS5beM= github.com/Masterminds/squirrel v1.5.4 h1:uUcX/aBc8O7Fg9kaISIUsHXdKuqehiXAMQTYX8afzqM= github.com/Masterminds/squirrel v1.5.4/go.mod h1:NNaOrjSoIDfDA40n7sr2tPNZRfjzjA400rg+riTZj10= -github.com/andybalholm/cascadia v1.3.4 h1:vM2lgh0Vru9Vwyfm4cQqWP2HHMW0u0+2PAW7Q38Qufg= -github.com/andybalholm/cascadia v1.3.4/go.mod h1:BLRmbRjpEtNKieZOCCvYj4RqN+KRA41GBe/5O+G93kM= +github.com/andybalholm/cascadia v1.3.5 h1:RLjq12WJy58dN6eCIQrz0bAGZkztHWsEPFxP53Y7Ms8= +github.com/andybalholm/cascadia v1.3.5/go.mod h1:BLRmbRjpEtNKieZOCCvYj4RqN+KRA41GBe/5O+G93kM= github.com/atombender/go-jsonschema v0.20.0 h1:AHg0LeI0HcjQ686ALwUNqVJjNRcSXpIR6U+wC2J0aFY= github.com/atombender/go-jsonschema v0.20.0/go.mod h1:ZmbuR11v2+cMM0PdP6ySxtyZEGFBmhgF4xa4J6Hdls8= github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuPk= 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= @@ -29,10 +29,10 @@ github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSs github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.1 h1:5RVFMOWjMyRy8cARdy79nAmgYw3hK/4HUq48LQ6Wwqo= github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.1/go.mod h1:ZXNYxsqcloTdSy/rNShjYzMhyjf0LaoftYK0p+A3h40= -github.com/deluan/go-taglib v0.0.0-20260905051825-df1d035571df h1:LdLQVAWVc6hCzqnrfVIEXOhP+r0iSit+EvsXwZDyL70= -github.com/deluan/go-taglib v0.0.0-20260905051825-df1d035571df/go.mod h1:QGxQ4Z1IWyY9w56xNEFjYAaWE8uSxA/gneQ7RPcFJrY= -github.com/deluan/rest v0.0.0-20211102003136-6260bc399cbf h1:tb246l2Zmpt/GpF9EcHCKTtwzrd0HGfEmoODFA/qnk4= -github.com/deluan/rest v0.0.0-20211102003136-6260bc399cbf/go.mod h1:tSgDythFsl0QgS/PFWfIZqcJKnkADWneY80jaVRlqK8= +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 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= @@ -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= @@ -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= @@ -88,31 +96,31 @@ 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= github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= -github.com/gohugoio/hashstructure v1.0.0 h1:vWYuyzs1n0LdI0F54TJQeYAiB44fHX7H9hCp9X6gHKg= -github.com/gohugoio/hashstructure v1.0.0/go.mod h1:FSbTK4QwxucJ2bC4Lvrs9a6x0DbQDXNoyBO+h4nlCgE= +github.com/gohugoio/hashstructure v1.1.0 h1:38yUfZBca6qXSbUpteLhjDGLNskclHaguFBYpjaRjf4= +github.com/gohugoio/hashstructure v1.1.0/go.mod h1:Pz8dcwjZs6FBKWu9x/ZIChrTHIM175zfUJK0KLvC1z8= github.com/golang/protobuf v1.3.1/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= 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= 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/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= @@ -128,17 +136,14 @@ 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= github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51/go.mod h1:CzGEWj7cYgsdH8dAjBGEr58BoE7ScuLd+fwFZ44+/x8= -github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= -github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/klauspost/compress v1.19.2 h1:hMRETovs/pu/dVWN7zIT1PGG8t509MwT6bO7XSi26R8= +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= @@ -159,16 +164,16 @@ github.com/lestrrat-go/httpcc v1.0.1 h1:ydWCStUeJLkpYyjLDHihupbn2tYmZ7m22BGkcvZZ github.com/lestrrat-go/httpcc v1.0.1/go.mod h1:qiltp3Mt56+55GPVCbTdM9MlqhvzyuL6W/NMDA8vA5E= github.com/lestrrat-go/httprc/v3 v3.0.6 h1:4FpLQ18KK/ypPbVU3NLWJNRvH3kcYiqKqWfKGqNWxxI= github.com/lestrrat-go/httprc/v3 v3.0.6/go.mod h1:mSMtkZW92Z98M5YoNNztbRGxbXHql7tSitCvaxvo9l0= -github.com/lestrrat-go/jwx/v3 v3.2.0 h1:Jb3zBASTSZXz7gzzSAfYqxXF8KejvKC4xWoePLQqXCA= -github.com/lestrrat-go/jwx/v3 v3.2.0/go.mod h1:38vQ8iWKq3qRSbilbzvzdQPuywhowwuR03lhkYskyrw= +github.com/lestrrat-go/jwx/v3 v3.3.0 h1:OXcYvQOQ7cxWzeZ/Q9sYk8ABe/kCSI371WmuACiCT+4= +github.com/lestrrat-go/jwx/v3 v3.3.0/go.mod h1:eIJhDcKHBwcgxqv8RiIylV67TVl1wJp/265IAHY1Db8= github.com/lestrrat-go/option/v2 v2.0.0 h1:XxrcaJESE1fokHy3FpaQ/cXW8ZsIdWcdFzzLOcID3Ss= github.com/lestrrat-go/option/v2 v2.0.0/go.mod h1:oSySsmzMoR0iRzCDCaUfsCzxQHUEuhOViQObyy7S6Vg= github.com/maruel/natural v1.3.0 h1:VsmCsBmEyrR46RomtgHs5hbKADGRVtliHTyCOLFBpsg= github.com/maruel/natural v1.3.0/go.mod h1:v+Rfd79xlw1AgVBjbO0BEQmptqb5HvL/k9GRHB7ZKEg= -github.com/mattn/go-isatty v0.0.23 h1:cYwCQTQf3HB6xUC+BtyCLZNr7IzbOmoZbmssVNzSyiQ= -github.com/mattn/go-isatty v0.0.23/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= -github.com/mattn/go-sqlite3 v1.14.50 h1:dmdFvo1XG4MPzA4IkAmE9upVz/Nj31uRoM5+jC8hYbY= -github.com/mattn/go-sqlite3 v1.14.50/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= +github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI= +github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A= +github.com/mattn/go-sqlite3 v1.14.52 h1:wVbm2Qnf4OXkqhBTSPuCRZDRnxfbVrrmiCEroVdog8U= +github.com/mattn/go-sqlite3 v1.14.52/go.mod h1:6JTjA44L93a0QCyJef5YvlPoKXntQPjzWv5gtm9sB6w= github.com/mfridman/interpolate v0.0.2 h1:pnuTK7MQIxxFz1Gr+rjSIx9u7qVjf5VOoM/u6BbAxPY= github.com/mfridman/interpolate v0.0.2/go.mod h1:p+7uk6oE07mpE/Ik1b8EckO0O4ZXiGAfshKBWLUM9Xg= github.com/mfridman/tparse v0.18.0 h1:wh6dzOKaIwkUGyKgOntDW4liXSo37qg5AXbIhkMV3vE= @@ -183,12 +188,16 @@ 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.32.1 h1:6tlvcDm/3sE8lGJbZ4+d4mO3RLy24/tQWOFzVSQNIfw= -github.com/onsi/ginkgo/v2 v2.32.1/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= @@ -199,16 +208,16 @@ github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZb github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pocketbase/dbx v1.12.0 h1:/oLErM+A0b4xI0PWTGPqSDVjzix48PqI/bng2l0PzoA= github.com/pocketbase/dbx v1.12.0/go.mod h1:xXRCIAKTHMgUCyCKZm55pUOdvFziJjQfXaWKhu2vhMs= -github.com/pressly/goose/v3 v3.27.3 h1:pIglVHjw99r4e/hDHHwbl9vfOsDMqUokfkXo6+n/RxA= -github.com/pressly/goose/v3 v3.27.3/go.mod h1:Dag+xpV6o20HR2LFY1j0q6MDwc3f7vPUFDA77R+0yGY= +github.com/pressly/goose/v3 v3.28.0 h1:D2M+iL31GmpZxSHOhX8mqyqAT3CXnokUmm0eKoSP+Vc= +github.com/pressly/goose/v3 v3.28.0/go.mod h1:v26MOuB8bL3kzzrt3Vqhb3R0PRVsl8hFQKdrht/L6Rk= github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU= github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE= github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE= github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi/PY= github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc= -github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI= -github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY= +github.com/prometheus/procfs v0.22.0 h1:6q9+/JL9IKAPbCmBrv9n5O5Ty3NKnciV5X7YGw0oics= +github.com/prometheus/procfs v0.22.0/go.mod h1:CvmFr/GVhIjIvWJZW3tgkODBQMRIf0EyWMQLHCHab58= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/rjeczalik/notify v0.9.3 h1:6rJAzHTGKXGj76sbRgDiDcYj/HniypXmSJo1SWakZeY= @@ -231,13 +240,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 +256,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= @@ -304,44 +307,41 @@ go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= -golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= -golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= -golang.org/x/image v0.45.0 h1:FMb1nTbH5H9vF55SriQHgFw5GnNL9Jg6L25BwXKzhB0= -golang.org/x/image v0.45.0/go.mod h1:n62x/7RqlwXDvGsSU4u6IUTUf6KghUZ9Bt7cG/T9Fx4= -golang.org/x/mod v0.40.0 h1:hUv+3cXcdRHz08UmSiOob7sadHig73uo5bkXxQ/tvUs= -golang.org/x/mod v0.40.0/go.mod h1:0/weTWkPWGBikyTWAX3dkjVztMmBA5hM0DH6BElSupE= -golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M= +golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA= +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-20190603091049-60506f45cf65/go.mod h1:HSz+uSET+XFnRR8LxR5pz3Of3rY3CfYBVs4xY44aLks= -golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= -golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= -golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= -golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues= +golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg= +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.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= -golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -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/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0= -golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w= +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-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= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= -golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= -golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= -golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= -golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= +golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI= +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= +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= 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= @@ -350,11 +350,11 @@ gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -modernc.org/libc v1.74.3 h1:a4J+Z8aVaxPyjyxRAdJzw246PqpcFGvVPnfT/AuM5Ws= -modernc.org/libc v1.74.3/go.mod h1:4H7h/MJ8wnjL8RAbp9v3OXgnk22X7MouHIhDbvP3gj4= +modernc.org/libc v1.75.6 h1:yKk8qo+Di4gkmvRboK8ocCqH22FiUCR6jRy2OwtCRus= +modernc.org/libc v1.75.6/go.mod h1:bO5o2ztHxBb2rjz0PgdHN0sSMw57CgxGFLZ3Qd/QpVQ= modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= -modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= -modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= -modernc.org/sqlite v1.54.0 h1:JCxR4qwkJvOaqAoYcgDoO25Nc+ROg6EJ2LfBVzdrgog= -modernc.org/sqlite v1.54.0/go.mod h1:4ntCLuNmnH8+GNqjka1wNg7KJd5/Hi5FYp8K+XQ7GZw= +modernc.org/memory v1.12.1 h1:nFMiWrpStgZczNl6XI9GnIk/rWhYIyHGUaR04pGbp9g= +modernc.org/memory v1.12.1/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/sqlite v1.57.0 h1:qNQP6xnx5M0ISNtlnxoOX0+cD5bJ0/gr9aMmndFczzg= +modernc.org/sqlite v1.57.0/go.mod h1:yCJ2cmAaIkHQ25oXWrF8H4O1lIfPYPR26yCEDj2P3pQ= 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"`)) + }) }) }) diff --git a/model/album.go b/model/album.go index ee24bfa96..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 { @@ -76,8 +77,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 } @@ -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/criteria/criteria.go b/model/criteria/criteria.go index 5d7dc3826..e6db3b481 100644 --- a/model/criteria/criteria.go +++ b/model/criteria/criteria.go @@ -64,18 +64,21 @@ func (c Criteria) IsPercentageLimit() bool { } func (c Criteria) ChildPlaylistIds() []string { - if c.Expression == nil { - return nil - } + return c.childPlaylistRefs(conjunction.ChildPlaylistIds) +} +func (c Criteria) ChildPlaylistPaths() []string { + return c.childPlaylistRefs(conjunction.ChildPlaylistPaths) +} + +func (c Criteria) childPlaylistRefs(extract func(conjunction) []string) []string { parent, ok := c.Expression.(conjunction) if !ok { return nil } - - ids := parent.ChildPlaylistIds() - slices.Sort(ids) - return slices.Compact(ids) + refs := extract(parent) + slices.Sort(refs) + return slices.Compact(refs) } func (c Criteria) MarshalJSON() ([]byte, error) { diff --git a/model/criteria/criteria_test.go b/model/criteria/criteria_test.go index 5e653150a..de5124568 100644 --- a/model/criteria/criteria_test.go +++ b/model/criteria/criteria_test.go @@ -323,19 +323,23 @@ var _ = Describe("Criteria", func() { Context("with child playlists", func() { var ( - topLevelInPlaylistID string - topLevelNotInPlaylistID string - nestedAnyInPlaylistID string - nestedAnyNotInPlaylistID string - nestedAllInPlaylistID string - nestedAllNotInPlaylistID string + topLevelInPlaylistID string + topLevelInPlaylistPath string + topLevelNotInPlaylistID string + nestedAnyInPlaylistID string + nestedAnyNotInPlaylistID string + nestedAllInPlaylistID string + nestedAllNotInPlaylistID string + nestedAnyNotInPlaylistPath string ) BeforeEach(func() { topLevelInPlaylistID = uuid.NewString() + topLevelInPlaylistPath = "./test.nsp" topLevelNotInPlaylistID = uuid.NewString() nestedAnyInPlaylistID = uuid.NewString() nestedAnyNotInPlaylistID = uuid.NewString() + nestedAnyNotInPlaylistPath = "../not-in-playlist.m3u" nestedAllInPlaylistID = uuid.NewString() nestedAllNotInPlaylistID = uuid.NewString() @@ -343,10 +347,12 @@ var _ = Describe("Criteria", func() { goObj = Criteria{ Expression: All{ InPlaylist{"id": topLevelInPlaylistID}, + InPlaylist{"path": topLevelInPlaylistPath}, NotInPlaylist{"id": topLevelNotInPlaylistID}, Any{ InPlaylist{"id": nestedAnyInPlaylistID}, NotInPlaylist{"id": nestedAnyNotInPlaylistID}, + NotInPlaylist{"path": nestedAnyNotInPlaylistPath}, }, All{ InPlaylist{"id": nestedAllInPlaylistID}, @@ -359,6 +365,18 @@ var _ = Describe("Criteria", func() { ids := goObj.ChildPlaylistIds() gomega.Expect(ids).To(gomega.ConsistOf(topLevelInPlaylistID, topLevelNotInPlaylistID, nestedAnyInPlaylistID, nestedAnyNotInPlaylistID, nestedAllInPlaylistID, nestedAllNotInPlaylistID)) }) + It("extracts all child smart playlist paths from expression criteria", func() { + paths := goObj.ChildPlaylistPaths() + gomega.Expect(paths).To(gomega.ConsistOf(topLevelInPlaylistPath, nestedAnyNotInPlaylistPath)) + }) + It("ignores empty child playlist paths", func() { + c := Criteria{Expression: All{InPlaylist{"path": ""}, NotInPlaylist{"path": ""}}} + gomega.Expect(c.ChildPlaylistPaths()).To(gomega.BeEmpty()) + }) + It("ignores empty child playlist ids", func() { + c := Criteria{Expression: All{InPlaylist{"id": ""}, NotInPlaylist{"id": ""}}} + gomega.Expect(c.ChildPlaylistIds()).To(gomega.BeEmpty()) + }) It("extracts child smart playlist IDs from deeply nested expression", func() { goObj = Criteria{ Expression: Any{ diff --git a/model/criteria/operators.go b/model/criteria/operators.go index 14a02ff4b..ec32f2d16 100644 --- a/model/criteria/operators.go +++ b/model/criteria/operators.go @@ -1,8 +1,9 @@ package criteria -// Conjunctions need to implement this interface, to allow Criteria to extract child playlist IDs recursively +// Conjunctions need to implement this interface, to allow Criteria to extract child playlist references recursively type conjunction interface { ChildPlaylistIds() []string + ChildPlaylistPaths() []string } type ( @@ -16,9 +17,9 @@ func (all All) MarshalJSON() ([]byte, error) { return marshalConjunction("all", all) } -func (all All) ChildPlaylistIds() (ids []string) { - return extractPlaylistIds(all) -} +func (all All) ChildPlaylistIds() []string { return extractPlaylistField(all, "id") } + +func (all All) ChildPlaylistPaths() []string { return extractPlaylistField(all, "path") } type ( Any []Expression @@ -31,9 +32,9 @@ func (any Any) MarshalJSON() ([]byte, error) { return marshalConjunction("any", any) } -func (any Any) ChildPlaylistIds() (ids []string) { - return extractPlaylistIds(any) -} +func (any Any) ChildPlaylistIds() []string { return extractPlaylistField(any, "id") } + +func (any Any) ChildPlaylistPaths() []string { return extractPlaylistField(any, "path") } type Is map[string]any type Eq = Is @@ -172,28 +173,20 @@ func (ip IsPresent) MarshalJSON() ([]byte, error) { func (ip IsPresent) fields() map[string]any { return ip } -func extractPlaylistIds(inputRule any) (ids []string) { - var id string - var ok bool - +func extractPlaylistField(inputRule any, field string) (values []string) { switch rule := inputRule.(type) { case Any: for _, rules := range rule { - ids = append(ids, extractPlaylistIds(rules)...) + values = append(values, extractPlaylistField(rules, field)...) } case All: for _, rules := range rule { - ids = append(ids, extractPlaylistIds(rules)...) + values = append(values, extractPlaylistField(rules, field)...) } - case InPlaylist: - if id, ok = rule["id"].(string); ok { - ids = append(ids, id) - } - case NotInPlaylist: - if id, ok = rule["id"].(string); ok { - ids = append(ids, id) + case InPlaylist, NotInPlaylist: + if value, ok := rule.(Expression).fields()[field].(string); ok && value != "" { + values = append(values, value) } } - return } diff --git a/model/datastore.go b/model/datastore.go index 273ca714b..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,36 +15,33 @@ 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 + // 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/model/errors.go b/model/errors.go index 9e56f9501..b3f07565f 100644 --- a/model/errors.go +++ b/model/errors.go @@ -1,11 +1,15 @@ package model -import "errors" +import ( + "errors" + + "github.com/deluan/rest" +) var ( - ErrNotFound = errors.New("data not found") + ErrNotFound = rest.ErrNotFound ErrInvalidAuth = errors.New("invalid authentication") - ErrNotAuthorized = errors.New("not authorized") + ErrNotAuthorized = rest.ErrPermissionDenied ErrExpired = errors.New("access expired") ErrNotAvailable = errors.New("functionality not available") ErrValidation = errors.New("validation error") 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/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..e80d22c89 100644 --- a/model/library.go +++ b/model/library.go @@ -1,8 +1,13 @@ package model import ( + "cmp" + "context" + "strings" "time" + "github.com/deluan/rest" + "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/utils/slice" ) @@ -25,6 +30,38 @@ type Library struct { TotalSize int64 `json:"totalSize" db:"total_size"` TotalDuration float64 `json:"totalDuration" db:"total_duration"` DefaultNewUsers bool `json:"defaultNewUsers" db:"default_new_users"` + PIDAlbum string `json:"pidAlbum" db:"pid_album"` + PIDTrack string `json:"pidTrack" db:"pid_track"` + ScannedPIDAlbum string `json:"-" db:"scanned_pid_album"` + ScannedPIDTrack string `json:"-" db:"scanned_pid_track"` +} + +// PIDConfig holds the persistent ID specs used to compute track and album IDs. +type PIDConfig struct { + Track string + Album string +} + +// EffectivePID returns the PID specs in effect for this library: its own overrides, falling back to +// the global config. +func (l Library) EffectivePID() PIDConfig { + return PIDConfig{ + Track: cmp.Or(l.PIDTrack, conf.Server.PID.Track), + Album: cmp.Or(l.PIDAlbum, conf.Server.PID.Album), + } +} + +// PIDChanged reports whether the effective PID specs differ from the ones used by the last finished +// scan of this library. A library that was never scanned counts as changed. +func (l Library) PIDChanged() bool { + pid := l.EffectivePID() + return !strings.EqualFold(l.ScannedPIDAlbum, pid.Album) || !strings.EqualFold(l.ScannedPIDTrack, pid.Track) +} + +// NeedsPIDRescan reports whether the library has content imported with an old PID config, so it must be +// rescanned in full. A library that never finished a scan has nothing to regroup. +func (l Library) NeedsPIDRescan() bool { + return !l.LastScanAt.IsZero() && l.PIDChanged() } const ( @@ -39,23 +76,26 @@ 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 + // SetScannedPID records the PID specs used by the last finished scan of the library + SetScannedPID(ctx context.Context, id int, pid PIDConfig) error + ScanInProgress(ctx context.Context) (bool, error) + RefreshStats(ctx context.Context, id int) error } diff --git a/model/library_matcher.go b/model/library_matcher.go new file mode 100644 index 000000000..83af96f9f --- /dev/null +++ b/model/library_matcher.go @@ -0,0 +1,57 @@ +package model + +import ( + "cmp" + "path/filepath" + "slices" + "strings" +) + +// LibraryMatcher finds the library that contains an absolute path. +type LibraryMatcher struct { + libraries Libraries + cleanedPaths []string +} + +// NewLibraryMatcher sorts the libraries longest path first, so /music-classical is checked before /music. +func NewLibraryMatcher(libs Libraries) *LibraryMatcher { + libs = slices.Clone(libs) + slices.SortFunc(libs, func(i, j Library) int { + return cmp.Compare(len(j.Path), len(i.Path)) + }) + cleanedPaths := make([]string, len(libs)) + for i, lib := range libs { + cleanedPaths[i] = filepath.Clean(lib.Path) + } + return &LibraryMatcher{libraries: libs, cleanedPaths: cleanedPaths} +} + +// FindLibrary returns the library whose path contains absolutePath. +func (lm *LibraryMatcher) FindLibrary(absolutePath string) (Library, bool) { + for i, libPath := range lm.cleanedPaths { + // A cleaned path only ends with a separator when it is a filesystem root + if strings.HasPrefix(absolutePath, libPath) && (len(absolutePath) == len(libPath) || + absolutePath[len(libPath)] == filepath.Separator || strings.HasSuffix(libPath, string(filepath.Separator))) { + return lm.libraries[i], true + } + } + return Library{}, false +} + +// LibraryRelativePath rebases an absolute path onto the library root, as the scanner's io/fs sees it +// (forward slashes). Relative paths, and absolute paths outside the library root, are returned unchanged. +func LibraryRelativePath(libPath, path string) string { + if !filepath.IsAbs(path) { + return path + } + // The library root may be relative (e.g. the default "./music"); it resolves against the same cwd + absLib, err := filepath.Abs(libPath) + if err != nil { + return path + } + rel, err := filepath.Rel(absLib, path) + if err != nil || !filepath.IsLocal(rel) { + return path + } + return filepath.ToSlash(rel) +} diff --git a/model/library_matcher_test.go b/model/library_matcher_test.go new file mode 100644 index 000000000..09e6f7e33 --- /dev/null +++ b/model/library_matcher_test.go @@ -0,0 +1,91 @@ +package model_test + +import ( + "os" + "path/filepath" + + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("LibraryMatcher", func() { + // Paths are written Unix-style and converted, so they use the OS separator, as filepath.Abs output does + find := func(libs model.Libraries, path string) int { + for i := range libs { + libs[i].Path = filepath.FromSlash(libs[i].Path) + } + lib, ok := model.NewLibraryMatcher(libs).FindLibrary(filepath.FromSlash(path)) + if !ok { + return 0 + } + return lib.ID + } + + DescribeTable("matches the longest library path", + func(libs model.Libraries, path string, expectedID int) { + Expect(find(libs, path)).To(Equal(expectedID)) + }, + Entry("nested library", model.Libraries{{ID: 1, Path: "/music"}, {ID: 2, Path: "/music-classical"}, {ID: 3, Path: "/music-classical/opera"}}, "/music-classical/opera/subdir/track.mp3", 3), + Entry("sibling with a shared prefix", model.Libraries{{ID: 1, Path: "/music"}, {ID: 2, Path: "/music-classical"}}, "/music-classical/track.mp3", 2), + Entry("shorter library", model.Libraries{{ID: 1, Path: "/music"}, {ID: 2, Path: "/music-classical"}}, "/music/track.mp3", 1), + Entry("exact library root", model.Libraries{{ID: 1, Path: "/music"}, {ID: 2, Path: "/music-classical"}}, "/music-classical", 2), + Entry("deeply nested libraries", model.Libraries{{ID: 1, Path: "/media"}, {ID: 2, Path: "/media/audio"}, {ID: 3, Path: "/media/audio/classical"}, {ID: 4, Path: "/media/audio/classical/baroque"}}, "/media/audio/classical/mozart/track.mp3", 3), + Entry("prefix that is not a path boundary", model.Libraries{{ID: 1, Path: "/a"}, {ID: 2, Path: "/ab"}, {ID: 3, Path: "/abc"}}, "/ab/file.mp3", 2), + Entry("special characters match literally", model.Libraries{{ID: 1, Path: "/music[test]"}, {ID: 2, Path: "/music(backup)"}}, "/music[test]/track.mp3", 1), + Entry("library path with a trailing slash", model.Libraries{{ID: 1, Path: "/music/"}}, "/music/track.mp3", 1), + Entry("library at the filesystem root", model.Libraries{{ID: 1, Path: "/"}}, "/music/track.mp3", 1), + Entry("nested library under a root library", model.Libraries{{ID: 1, Path: "/"}, {ID: 2, Path: "/music"}}, "/music/track.mp3", 2), + ) + + It("does not match a path outside every library", func() { + Expect(find(model.Libraries{{ID: 1, Path: "/music"}}, "/music-backup/track.mp3")).To(BeZero()) + }) + + It("does not match anything without libraries", func() { + Expect(find(nil, "/music/track.mp3")).To(BeZero()) + }) + + It("does not reorder the caller's libraries", func() { + libs := model.Libraries{{ID: 1, Path: "/a"}, {ID: 2, Path: "/abc"}} + model.NewLibraryMatcher(libs) + Expect(libs.IDs()).To(Equal([]int{1, 2})) + }) +}) + +var _ = Describe("LibraryRelativePath", func() { + // Paths are built with filepath so the "absolute" cases stay absolute on every OS + // (a Unix-style "/foo" is not absolute on Windows). + libRoot, _ := filepath.Abs(filepath.Join("jukebox", "collection")) + outside, _ := filepath.Abs(filepath.Join("somewhere", "else")) + + It("returns a relative path unchanged", func() { + Expect(model.LibraryRelativePath(libRoot, "_Collection")).To(Equal("_Collection")) + }) + + It("rebases an absolute target when the library root is relative", func() { + cwd, err := os.Getwd() + Expect(err).ToNot(HaveOccurred()) + Expect(model.LibraryRelativePath(filepath.Join("music", "library"), filepath.Join(cwd, "music", "library", "rock"))).To(Equal("rock")) + }) + + It("rebases an absolute path that equals the library root to '.'", func() { + Expect(model.LibraryRelativePath(libRoot, libRoot)).To(Equal(".")) + }) + + It("rebases an absolute path under the library root", func() { + Expect(model.LibraryRelativePath(libRoot, filepath.Join(libRoot, "_Collection"))).To(Equal("_Collection")) + }) + + It("handles a trailing slash on the library path", func() { + Expect(model.LibraryRelativePath(libRoot+string(filepath.Separator), filepath.Join(libRoot, "_Collection"))).To(Equal("_Collection")) + }) + + It("leaves an absolute path outside the library root unchanged", func() { + Expect(model.LibraryRelativePath(libRoot, outside)).To(Equal(outside)) + }) + + It("returns an empty path unchanged", func() { + Expect(model.LibraryRelativePath(libRoot, "")).To(Equal("")) + }) +}) diff --git a/model/library_test.go b/model/library_test.go new file mode 100644 index 000000000..4e799e13f --- /dev/null +++ b/model/library_test.go @@ -0,0 +1,73 @@ +package model_test + +import ( + "encoding/json" + "time" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Library PID config", func() { + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.PID.Album = "global_album" + conf.Server.PID.Track = "global_track" + }) + + Describe("EffectivePID", func() { + It("falls back to the global config", func() { + Expect(model.Library{}.EffectivePID()).To(Equal(model.PIDConfig{Track: "global_track", Album: "global_album"})) + }) + It("uses the library overrides", func() { + lib := model.Library{PIDAlbum: "folder", PIDTrack: "title"} + Expect(lib.EffectivePID()).To(Equal(model.PIDConfig{Track: "title", Album: "folder"})) + }) + }) + + Describe("PIDChanged", func() { + It("is false when the scanned specs match, ignoring case", func() { + lib := model.Library{ScannedPIDAlbum: "GLOBAL_ALBUM", ScannedPIDTrack: "global_track"} + Expect(lib.PIDChanged()).To(BeFalse()) + }) + It("is true when the album override differs from the scanned spec", func() { + lib := model.Library{PIDAlbum: "folder", ScannedPIDAlbum: "global_album", ScannedPIDTrack: "global_track"} + Expect(lib.PIDChanged()).To(BeTrue()) + }) + It("is true when only the track spec changed", func() { + lib := model.Library{PIDTrack: "title", ScannedPIDAlbum: "global_album", ScannedPIDTrack: "global_track"} + Expect(lib.PIDChanged()).To(BeTrue()) + }) + It("is true when the global config changed for a library without overrides", func() { + lib := model.Library{ScannedPIDAlbum: "old_album", ScannedPIDTrack: "global_track"} + Expect(lib.PIDChanged()).To(BeTrue()) + }) + It("is true for a library that was never scanned", func() { + Expect(model.Library{}.PIDChanged()).To(BeTrue()) + }) + }) + + Describe("NeedsPIDRescan", func() { + It("is false for a library that never finished a scan", func() { + Expect(model.Library{PIDAlbum: "folder"}.NeedsPIDRescan()).To(BeFalse()) + }) + It("is true for a scanned library whose PID config changed", func() { + lib := model.Library{PIDAlbum: "folder", ScannedPIDAlbum: "global_album", ScannedPIDTrack: "global_track", LastScanAt: time.Now()} + Expect(lib.NeedsPIDRescan()).To(BeTrue()) + }) + It("is false for a scanned library whose PID config did not change", func() { + lib := model.Library{ScannedPIDAlbum: "global_album", ScannedPIDTrack: "global_track", LastScanAt: time.Now()} + Expect(lib.NeedsPIDRescan()).To(BeFalse()) + }) + }) + + It("does not expose the scanned specs in JSON", func() { + data, err := json.Marshal(model.Library{PIDAlbum: "folder", ScannedPIDAlbum: "secret_album", ScannedPIDTrack: "secret_track"}) + Expect(err).ToNot(HaveOccurred()) + Expect(string(data)).To(ContainSubstring(`"pidAlbum":"folder"`)) + Expect(string(data)).ToNot(ContainSubstring("secret_")) + }) +}) diff --git a/model/mediafile.go b/model/mediafile.go index 2669018f3..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" @@ -105,15 +107,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 } @@ -537,39 +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(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/metadata/map_mediafile.go b/model/metadata/map_mediafile.go index b3ce4ef02..2135824d2 100644 --- a/model/metadata/map_mediafile.go +++ b/model/metadata/map_mediafile.go @@ -8,15 +8,14 @@ import ( "math" "strconv" - "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils/str" ) -func (md Metadata) ToMediaFile(libID int, folderID string) model.MediaFile { +func (md Metadata) ToMediaFile(lib model.Library, folderID string) model.MediaFile { mf := model.MediaFile{ - LibraryID: libID, + LibraryID: lib.ID, FolderID: folderID, Tags: maps.Clone(md.tags), } @@ -37,8 +36,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() @@ -84,8 +83,9 @@ func (md Metadata) ToMediaFile(libID int, folderID string) model.MediaFile { mf.AlbumArtist = md.mapDisplayAlbumArtist(mf) // Persistent IDs - mf.PID = md.trackPID(mf) - mf.AlbumID = md.albumID(mf, conf.Server.PID.Album) + pid := lib.EffectivePID() + mf.PID = md.trackPID(mf, pid) + mf.AlbumID = md.albumID(mf, pid.Album) // BFR These IDs will go away once the UI handle multiple participants. // BFR For Legacy Subsonic compatibility, we will set them in the API handlers diff --git a/model/metadata/map_mediafile_test.go b/model/metadata/map_mediafile_test.go index 75a7ed358..c19398841 100644 --- a/model/metadata/map_mediafile_test.go +++ b/model/metadata/map_mediafile_test.go @@ -30,9 +30,23 @@ var _ = Describe("ToMediaFile", func() { var toMediaFile = func(tags model.RawTags) model.MediaFile { props.Tags = tags md = metadata.New("filepath", props) - return md.ToMediaFile(1, "folderID") + return md.ToMediaFile(model.Library{ID: 1}, "folderID") } + Describe("Persistent IDs", func() { + It("uses the library PID config for the album ID and for albumid in the track spec", func() { + props.Tags = model.RawTags{"ALBUM": {"Kind of Blue"}, "TITLE": {"So What"}} + md = metadata.New("Jazz/Loose/01.mp3", props) + + byTags := md.ToMediaFile(model.Library{ID: 1, PIDAlbum: "album", PIDTrack: "albumid,title"}, "folderID") + byFolder := md.ToMediaFile(model.Library{ID: 1, PIDAlbum: "folder", PIDTrack: "albumid,title"}, "folderID") + + Expect(byFolder.AlbumID).ToNot(Equal(byTags.AlbumID)) + Expect(byFolder.AlbumID).To(Equal(md.AlbumID(byFolder, "folder"))) + Expect(byFolder.PID).ToNot(Equal(byTags.PID)) + }) + }) + Describe("Dates", func() { It("should parse properly tagged dates ", func() { mf = toMediaFile(model.RawTags{ @@ -131,6 +145,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/map_participants_test.go b/model/metadata/map_participants_test.go index ec66e12b9..db652fb8b 100644 --- a/model/metadata/map_participants_test.go +++ b/model/metadata/map_participants_test.go @@ -38,7 +38,7 @@ var _ = Describe("Participants", func() { var toMediaFile = func(tags model.RawTags) model.MediaFile { props.Tags = tags md = metadata.New("filepath", props) - return md.ToMediaFile(1, "folderID") + return md.ToMediaFile(model.Library{ID: 1}, "folderID") } Describe("ARTIST(S) tags", 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..a1a675006 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() { @@ -319,7 +323,7 @@ var _ = Describe("Metadata", func() { tag: {tagValue}, } md = metadata.New(filePath, props) - return md.ToMediaFile(0, "0") + return md.ToMediaFile(model.Library{}, "0") } DescribeTable("Gain", diff --git a/model/metadata/persistent_ids.go b/model/metadata/persistent_ids.go index db315dc6b..b66ce824a 100644 --- a/model/metadata/persistent_ids.go +++ b/model/metadata/persistent_ids.go @@ -2,11 +2,11 @@ package metadata import ( "cmp" + "errors" "fmt" "path/filepath" "strings" - "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" @@ -22,12 +22,13 @@ type hashFunc = func(...string) string // attributes. Attributes can be either tags or processed values like folder, // albumid, albumartistid, etc. For each field, it gets all its attribute values // and concatenates them, then hashes the result. If a field is empty, it is -// skipped and the function looks for the next field. +// skipped and the function looks for the next field. albumSpec is the album PID +// spec used to resolve the `albumid` attribute. // // Taking hash as a parameter (instead of closing over it in a factory) keeps // mf on the stack: closing over mf would force the whole ~1KB MediaFile to the // heap on every call. -func computePID(mf model.MediaFile, md Metadata, spec string, prependLibId bool, hash hashFunc) string { +func computePID(mf model.MediaFile, md Metadata, spec, albumSpec string, prependLibId bool, hash hashFunc) string { switch spec { case "track_legacy": return legacyTrackID(mf, prependLibId) @@ -41,7 +42,7 @@ func computePID(mf model.MediaFile, md Metadata, spec string, prependLibId bool, values := make([]string, len(attributes)) hasValue := false for i, attr := range attributes { - v := getPIDAttr(mf, md, attr, prependLibId, spec, hash) + v := getPIDAttr(mf, md, attr, prependLibId, spec, albumSpec, hash) if v != "" { hasValue = true } @@ -58,15 +59,15 @@ func computePID(mf model.MediaFile, md Metadata, spec string, prependLibId bool, return hash(pid) } -func getPIDAttr(mf model.MediaFile, md Metadata, attr string, prependLibId bool, spec string, hash hashFunc) string { +func getPIDAttr(mf model.MediaFile, md Metadata, attr string, prependLibId bool, spec, albumSpec string, hash hashFunc) string { attr = strings.TrimSpace(strings.ToLower(attr)) switch attr { case "albumid": - if spec == conf.Server.PID.Album { + if spec == albumSpec { log.Error("Recursive PID definition detected, ignoring `albumid`", "spec", spec) return "" } - return computePID(mf, md, conf.Server.PID.Album, prependLibId, hash) + return computePID(mf, md, albumSpec, albumSpec, prependLibId, hash) case "folder": return filepath.Dir(mf.Path) case "albumartistid": @@ -79,18 +80,50 @@ func getPIDAttr(mf model.MediaFile, md Metadata, attr string, prependLibId bool, return md.String(model.TagName(attr)) } -func (md Metadata) trackPID(mf model.MediaFile) string { - return computePID(mf, md, conf.Server.PID.Track, true, id.NewHash) +// ValidatePIDSpec checks a PID override before it is stored; empty means "use the global config". +// Aliases resolve to empty at scan time: accepted only in track specs, because the default one uses them. +func ValidatePIDSpec(spec string, isAlbum bool) error { + switch { + case spec == "", isAlbum && spec == "album_legacy", !isAlbum && spec == "track_legacy": + return nil + } + for field := range strings.SplitSeq(spec, "|") { + for attr := range strings.SplitSeq(field, ",") { + attr = strings.TrimSpace(strings.ToLower(attr)) + switch attr { + case "": + return fmt.Errorf("empty attribute in %q", spec) + case "albumid": + if isAlbum { + return errors.New("albumid cannot be used in an album PID") + } + case "folder", "albumartistid": + default: + name, ok := model.CanonicalTagName(attr) + if !ok { + return fmt.Errorf("unknown attribute %q", attr) + } + if isAlbum && string(name) != attr { + return fmt.Errorf("use the tag name %q instead of its alias %q", name, attr) + } + } + } + } + return nil +} + +func (md Metadata) trackPID(mf model.MediaFile, pid model.PIDConfig) string { + return computePID(mf, md, pid.Track, pid.Album, true, id.NewHash) } func (md Metadata) albumID(mf model.MediaFile, pidConf string) string { - return computePID(mf, md, pidConf, true, id.NewHash) + return computePID(mf, md, pidConf, pidConf, true, id.NewHash) } // BFR Must be configurable? func (md Metadata) artistID(name string) string { mf := model.MediaFile{AlbumArtist: name} - return computePID(mf, md, "albumartistid", false, id.NewHash) + return computePID(mf, md, "albumartistid", "", false, id.NewHash) } func (md Metadata) mapTrackTitle() string { diff --git a/model/metadata/persistent_ids_test.go b/model/metadata/persistent_ids_test.go index 8e38bbd42..9f6eaf1f4 100644 --- a/model/metadata/persistent_ids_test.go +++ b/model/metadata/persistent_ids_test.go @@ -3,8 +3,7 @@ package metadata import ( "strings" - "github.com/navidrome/navidrome/conf" - "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" @@ -13,16 +12,18 @@ import ( var _ = Describe("getPID", func() { var ( - md Metadata - mf model.MediaFile - sum hashFunc + md Metadata + mf model.MediaFile + sum hashFunc + albumSpec string ) getPID := func(mf model.MediaFile, md Metadata, spec string, prependLibId bool) string { - return computePID(mf, md, spec, prependLibId, sum) + return computePID(mf, md, spec, albumSpec, prependLibId, sum) } BeforeEach(func() { sum = func(s ...string) string { return "(" + strings.Join(s, ",") + ")" } + albumSpec = consts.DefaultAlbumPID }) Context("attributes are tags", func() { @@ -66,8 +67,7 @@ var _ = Describe("getPID", func() { Context("calculated attributes", func() { BeforeEach(func() { - DeferCleanup(configtest.SetupConfig()) - conf.Server.PID.Album = "musicbrainz_albumid|albumartistid,album,albumversion,releasedate" + albumSpec = "musicbrainz_albumid|albumartistid,album,albumversion,releasedate" }) When("field is title", func() { It("should return the pid", func() { @@ -121,8 +121,8 @@ var _ = Describe("getPID", func() { When("albumid configuration refers to albumid recursively", func() { It("should avoid infinite recursion", func() { // Reproduce the issue from #4920 - conf.Server.PID.Album = "albumid,album,albumversion,releasedate" - spec := conf.Server.PID.Album + albumSpec = "albumid,album,albumversion,releasedate" + spec := albumSpec md.tags = map[model.TagName][]string{ "album": {"Album Name"}, "albumversion": {"Version"}, @@ -205,8 +205,7 @@ var _ = Describe("getPID", func() { }) When("prependLibId is true with nested albumid", func() { It("should handle nested albumid calls correctly", func() { - DeferCleanup(configtest.SetupConfig()) - conf.Server.PID.Album = "album" + albumSpec = "album" spec := "albumid" md.tags = map[model.TagName][]string{"album": {"Test Album"}} mf.AlbumArtist = "Test Artist" @@ -306,3 +305,34 @@ var _ = Describe("getPID", func() { }) }) }) + +var _ = Describe("ValidatePIDSpec", func() { + DescribeTable("accepts valid specs", + func(spec string, isAlbum bool) { + Expect(ValidatePIDSpec(spec, isAlbum)).To(Succeed()) + }, + Entry("empty, meaning the global config", "", true), + Entry("default album spec", consts.DefaultAlbumPID, true), + Entry("default track spec, which uses tag aliases", consts.DefaultTrackPID, false), + Entry("folder", "folder", true), + Entry("album legacy", "album_legacy", true), + Entry("track legacy", "track_legacy", false), + Entry("computed attributes", "albumartistid,album|title", true), + Entry("albumid in a track spec", "albumid,title", false), + Entry("spaces and mixed case", "MusicBrainz_AlbumID | Folder", true), + ) + + DescribeTable("rejects invalid specs", + func(spec string, isAlbum bool, msg string) { + Expect(ValidatePIDSpec(spec, isAlbum)).To(MatchError(ContainSubstring(msg))) + }, + Entry("unknown tag", "albmversion", true, `unknown attribute "albmversion"`), + Entry("empty field", "album||title", true, "empty attribute"), + Entry("empty attribute", "album,,title", true, "empty attribute"), + Entry("trailing separator", "album|", true, "empty attribute"), + Entry("albumid in an album spec", "albumid,album", true, "albumid"), + Entry("tag alias in an album spec", "talb", true, `use the tag name "album" instead of its alias "talb"`), + Entry("track legacy in an album spec", "track_legacy", true, `unknown attribute "track_legacy"`), + Entry("album legacy in a track spec", "album_legacy", false, `unknown attribute "album_legacy"`), + ) +}) diff --git a/model/player.go b/model/player.go index 39ea99d1a..c03058419 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 { @@ -18,14 +21,20 @@ 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 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) + FindByAPIKey(ctx context.Context, key string) (*Player, error) + SetAPIKey(ctx context.Context, playerID, key string) error } diff --git a/model/playlist.go b/model/playlist.go index abd11b8c4..d2ed97682 100644 --- a/model/playlist.go +++ b/model/playlist.go @@ -1,13 +1,19 @@ package model import ( + "context" "iter" + "maps" + "os" + "path/filepath" "slices" "strconv" "time" + "github.com/deluan/rest" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model/criteria" ) @@ -137,24 +143,84 @@ func (pls Playlist) UploadedImagePath() string { return UploadedImagePath(consts.EntityPlaylist, pls.UploadedImage) } +// NormalizedRules returns the rules with child playlist paths resolved to absolute, OS-native paths. +func (pls Playlist) NormalizedRules() *criteria.Criteria { + if pls.Rules == nil || pls.Rules.Expression == nil { + return pls.Rules + } + + rules := *pls.Rules + rules.Expression = normalizePlaylistPaths(pls.Rules.Expression, pls.Path) + return &rules +} + +func normalizePlaylistPaths(inputRule criteria.Expression, referencingPlaylistPath string) criteria.Expression { + switch rule := inputRule.(type) { + case criteria.Any: + anyCriteria := make(criteria.Any, len(rule)) + for i, rules := range rule { + anyCriteria[i] = normalizePlaylistPaths(rules, referencingPlaylistPath) + } + return anyCriteria + case criteria.All: + allCriteria := make(criteria.All, len(rule)) + for i, rules := range rule { + allCriteria[i] = normalizePlaylistPaths(rules, referencingPlaylistPath) + } + return allCriteria + case criteria.InPlaylist: + return criteria.InPlaylist(normalizeChildPathRule(rule, referencingPlaylistPath)) + case criteria.NotInPlaylist: + return criteria.NotInPlaylist(normalizeChildPathRule(rule, referencingPlaylistPath)) + } + + return inputRule +} + +func normalizeChildPathRule(rule map[string]any, referencingPlaylistPath string) map[string]any { + path, ok := rule["path"].(string) + if !ok || path == "" { + return rule + } + + // References use forward slashes to stay portable, while Playlist.Path is OS-native. + path = filepath.FromSlash(path) + switch { + case isAbsPlaylistRef(path): + path = filepath.Clean(path) + case referencingPlaylistPath != "": + path = filepath.Join(filepath.Dir(referencingPlaylistPath), path) + default: + log.Warn("Cannot resolve relative playlist reference: playlist has no file path", "reference", path) + } + normalized := maps.Clone(rule) + normalized["path"] = path + return normalized +} + +// filepath.IsAbs rejects a bare leading separator on Windows, but that is how Unix spells absolute. +func isAbsPlaylistRef(path string) bool { + return filepath.IsAbs(path) || os.IsPathSeparator(path[0]) +} + 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 { @@ -177,17 +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) - 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/playlist_test.go b/model/playlist_test.go index d98c85716..7ffb69b63 100644 --- a/model/playlist_test.go +++ b/model/playlist_test.go @@ -1,6 +1,7 @@ package model_test import ( + "path/filepath" "time" "github.com/navidrome/navidrome/conf" @@ -88,4 +89,102 @@ var _ = Describe("Playlist", func() { Expect(model.Playlist{Sync: true}.TracksEditable()).To(BeFalse()) }) }) + + Describe("NormalizedRules()", func() { + // absPath builds an OS-native absolute path so these specs also run on Windows. + absPath := func(parts ...string) string { + abs, err := filepath.Abs(filepath.Join(parts...)) + Expect(err).ToNot(HaveOccurred()) + return abs + } + normalize := func(pls model.Playlist) criteria.Expression { + return pls.NormalizedRules().Expression + } + + It("resolves relative references against the playlist folder", func() { + pls := model.Playlist{ + Path: absPath("test", "nested", "my-playlist.nsp"), + Rules: &criteria.Criteria{Expression: criteria.All{ + criteria.InPlaylist{"path": "../up.m3u"}, + criteria.NotInPlaylist{"path": "./sibling.nsp"}, + criteria.Any{criteria.InPlaylist{"path": "sub/deep.nsp"}}, + }}, + } + Expect(normalize(pls)).To(BeEquivalentTo(criteria.All{ + criteria.InPlaylist{"path": absPath("test", "up.m3u")}, + criteria.NotInPlaylist{"path": absPath("test", "nested", "sibling.nsp")}, + criteria.Any{criteria.InPlaylist{"path": absPath("test", "nested", "sub", "deep.nsp")}}, + })) + }) + + It("cleans absolute references", func() { + dirty := absPath("music") + string(filepath.Separator) + "." + string(filepath.Separator) + "child.nsp" + pls := model.Playlist{ + Path: absPath("test", "my-playlist.nsp"), + Rules: &criteria.Criteria{Expression: criteria.All{criteria.NotInPlaylist{"path": dirty}}}, + } + Expect(normalize(pls)).To(BeEquivalentTo(criteria.All{ + criteria.NotInPlaylist{"path": absPath("music", "child.nsp")}, + })) + }) + + It("treats a leading slash as absolute on every OS", func() { + pls := model.Playlist{ + Path: absPath("test", "my-playlist.nsp"), + Rules: &criteria.Criteria{Expression: criteria.All{criteria.InPlaylist{"path": "/other/./root.m3u"}}}, + } + Expect(normalize(pls)).To(BeEquivalentTo(criteria.All{ + criteria.InPlaylist{"path": filepath.FromSlash("/other/root.m3u")}, + })) + }) + + It("leaves empty paths and id references untouched", func() { + pls := model.Playlist{ + Path: absPath("test", "my-playlist.nsp"), + Rules: &criteria.Criteria{Expression: criteria.All{ + criteria.InPlaylist{"path": ""}, + criteria.InPlaylist{"id": "94d8ba52-7aca-40e2-af82-4cb09c43d710"}, + criteria.Eq{"artist": "Bob Dealin"}, + }}, + } + Expect(normalize(pls)).To(BeEquivalentTo(criteria.All{ + criteria.InPlaylist{"path": ""}, + criteria.InPlaylist{"id": "94d8ba52-7aca-40e2-af82-4cb09c43d710"}, + criteria.Eq{"artist": "Bob Dealin"}, + })) + }) + + It("skips relative references when the playlist has no path", func() { + pls := model.Playlist{ + Rules: &criteria.Criteria{Expression: criteria.All{criteria.InPlaylist{"path": "../up.m3u"}}}, + } + Expect(normalize(pls)).To(BeEquivalentTo(criteria.All{ + criteria.InPlaylist{"path": filepath.FromSlash("../up.m3u")}, + })) + }) + + It("preserves every other criteria field", func() { + rules := criteria.Criteria{ + Expression: criteria.All{criteria.InPlaylist{"path": "child.nsp"}}, + Sort: "title", + Order: "desc", + Limit: 10, + LimitPercent: 25, + Offset: 5, + RefreshDelay: 3 * time.Hour, + } + pls := model.Playlist{Path: absPath("test", "my-playlist.nsp"), Rules: &rules} + + normalized := *pls.NormalizedRules() + normalized.Expression = rules.Expression + Expect(normalized).To(Equal(rules)) + }) + + It("does not mutate the original playlist rules", func() { + original := criteria.All{criteria.InPlaylist{"path": "child.nsp"}} + pls := model.Playlist{Path: absPath("test", "my-playlist.nsp"), Rules: &criteria.Criteria{Expression: original}} + _ = pls.NormalizedRules() + Expect(original[0]).To(BeEquivalentTo(criteria.InPlaylist{"path": "child.nsp"})) + }) + }) }) 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/scanner.go b/model/scanner.go index 36c9007fb..d22c3d0d6 100644 --- a/model/scanner.go +++ b/model/scanner.go @@ -2,12 +2,15 @@ package model import ( "context" + "errors" "fmt" "strconv" "strings" "time" ) +var ErrAlreadyScanning = errors.New("already scanning") + // ScanTarget represents a specific folder within a library to be scanned. // NOTE: This struct is used as a map key, so it should only contain comparable types. type ScanTarget struct { 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 ce0846d60..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" ) @@ -32,9 +34,6 @@ type Share struct { func (s Share) CoverArtID() ArtworkID { ids := strings.SplitN(s.ResourceIDs, ",", 2) - if len(ids) == 0 { - return ArtworkID{} - } switch s.ResourceType { case "album": return Album{ID: ids[0]}.CoverArtID() @@ -43,6 +42,10 @@ func (s Share) CoverArtID() ArtworkID { case "artist": return Artist{ID: ids[0]}.CoverArtID() } + // Tracks can be empty when they went missing or the owner lost access to their library. + if len(s.Tracks) == 0 { + return ArtworkID{} + } rnd := random.Int64N(len(s.Tracks)) return s.Tracks[rnd].CoverArtID() } @@ -55,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/share_test.go b/model/share_test.go new file mode 100644 index 000000000..b5e38af59 --- /dev/null +++ b/model/share_test.go @@ -0,0 +1,19 @@ +package model_test + +import ( + "github.com/navidrome/navidrome/model" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Share.CoverArtID", func() { + It("returns an empty artwork ID for a media file share with no visible tracks", func() { + s := model.Share{ResourceType: "media_file", ResourceIDs: "mf-1"} + Expect(s.CoverArtID()).To(Equal(model.ArtworkID{})) + }) + + It("picks a track's cover for a media file share", func() { + s := model.Share{ResourceType: "media_file", ResourceIDs: "mf-1", Tracks: model.MediaFiles{{ID: "mf-1"}}} + Expect(s.CoverArtID()).To(Equal(model.MediaFile{ID: "mf-1"}.CoverArtID())) + }) +}) diff --git a/model/tag.go b/model/tag.go index 234cfb359..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" @@ -80,6 +82,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 { @@ -155,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/tag_mappings.go b/model/tag_mappings.go index ce7d2f37b..5a8168754 100644 --- a/model/tag_mappings.go +++ b/model/tag_mappings.go @@ -195,6 +195,28 @@ func TagMappings() map[TagName]TagConf { return mappings } +// CanonicalTagName returns the mapped tag that name is, or is an alias of. Tags are stored under this name. +func CanonicalTagName(name string) (TagName, bool) { + tagName, ok := tagNameIndex()[TagName(name).ToLower()] + return tagName, ok +} + +// tagNameIndex maps every tag name and alias to its tag name. Names are added last, so they win over aliases +// (musicbrainz_trackid is a tag and also an alias of musicbrainz_recordingid). +var tagNameIndex = sync.OnceValue(func() map[TagName]TagName { + mappings := TagMappings() + index := make(map[TagName]TagName, len(mappings)) + for name, tag := range mappings { + for _, alias := range tag.Aliases { + index[TagName(alias)] = name + } + } + for name := range mappings { + index[name] = name + } + return index +}) + func TagRolesConf() TagConf { _, cfg := parseMappings() return cfg.Roles diff --git a/model/tag_mappings_test.go b/model/tag_mappings_test.go index e582c3f2f..91e54e5d4 100644 --- a/model/tag_mappings_test.go +++ b/model/tag_mappings_test.go @@ -192,3 +192,22 @@ var _ = Describe("TagConf", func() { }) }) }) + +var _ = Describe("CanonicalTagName", func() { + DescribeTable("resolves tag names and aliases", + func(name string, expected TagName) { + tagName, ok := CanonicalTagName(name) + Expect(ok).To(BeTrue()) + Expect(tagName).To(Equal(expected)) + }, + Entry("tag name", "album", TagAlbum), + Entry("alias", "talb", TagAlbum), + Entry("mixed case alias", "TALB", TagAlbum), + Entry("tag name that is also an alias of another tag", "musicbrainz_trackid", TagMusicBrainzTrackID), + ) + + It("does not resolve an unknown name", func() { + _, ok := CanonicalTagName("nosuchtag") + Expect(ok).To(BeFalse()) + }) +}) 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/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..6a940a320 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()) @@ -132,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, @@ -190,49 +190,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 +242,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 +282,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 +295,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 +303,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 +351,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 +389,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 +408,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 +424,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 +449,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..0235805c0 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,17 +406,91 @@ 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)) } }) }) + + 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() { @@ -425,12 +500,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 +523,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 +545,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 +786,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 +802,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 +815,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 +836,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 +847,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 +866,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 +908,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 +931,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 +957,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 +971,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 +985,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 +1012,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 +1031,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 +1063,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 +1071,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 +1085,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 +1105,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 +1122,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 +1134,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 +1156,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 +1176,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 +1187,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 +1241,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 +1257,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 +1268,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 1ff291a0f..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 @@ -162,21 +162,34 @@ func NewArtistRepository(ctx context.Context, db dbx.Builder) model.ArtistReposi func roleFilter(_ string, role any) Sqlizer { if role, ok := role.(string); ok { - if _, ok := model.AllRoles[role]; ok { - return Expr("JSON_EXTRACT(library_artist.stats, '$." + role + ".m') IS NOT NULL") + if safe, ok := sanitizeArtistStatsRole(role); ok && safe != "total" { + return Expr("JSON_EXTRACT(library_artist.stats, '$." + safe + ".m') IS NOT NULL") } } return Eq{"1": 2} } +// sanitizeArtistStatsRole allowlists values interpolated into JSON paths for artist +// stats (filter and sort). "total" is the aggregate key stored by the scanner. +// Unknown values must not reach SQL string concatenation. +func sanitizeArtistStatsRole(role string) (string, bool) { + if role == "" || role == "total" { + return "total", true + } + if _, ok := model.AllRoles[role]; ok { + return role, true + } + return "", false +} + // artistLibraryIdFilter filters artists based on library access through the library_artist table func artistLibraryIdFilter(_ string, value any) Sqlizer { return Eq{"library_artist.library_id": value} } // 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") @@ -191,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 } @@ -311,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 @@ -339,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 } @@ -354,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 @@ -362,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 { @@ -381,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 @@ -394,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) @@ -405,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 @@ -416,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) } @@ -425,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 @@ -443,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 @@ -457,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 := ` @@ -468,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 } @@ -545,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)) @@ -571,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) } @@ -580,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 } @@ -636,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 @@ -704,34 +717,37 @@ 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 { - role = v + if safe, ok := sanitizeArtistStatsRole(v); ok { + role = safe + } } } - 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 d337b4c22..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() { @@ -112,13 +118,98 @@ var _ = Describe("ArtistRepository", func() { }) }) + Describe("sanitizeArtistStatsRole", func() { + It("allowlists total and registered roles", func() { + role, ok := sanitizeArtistStatsRole("total") + Expect(ok).To(BeTrue()) + Expect(role).To(Equal("total")) + role, ok = sanitizeArtistStatsRole("albumartist") + Expect(ok).To(BeTrue()) + Expect(role).To(Equal("albumartist")) + }) + + It("rejects SQL injection payloads used in sort mappings", func() { + payload := "total'||(SELECT password FROM user LIMIT 1)||'" + role, ok := sanitizeArtistStatsRole(payload) + Expect(ok).To(BeFalse()) + Expect(role).To(BeEmpty()) + }) + }) + + 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(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(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()) + }) + }) + Describe("searchScope", func() { // Resolves the library IDs a search must be restricted to (nil = fast-path / no filter), // 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}}} @@ -139,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) @@ -155,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()) @@ -264,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 }))) }) @@ -299,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()) @@ -312,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()) @@ -322,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)) }) @@ -346,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")) @@ -367,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")) @@ -398,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")) @@ -419,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")) @@ -452,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")) @@ -468,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")) @@ -486,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)) @@ -507,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)) }) @@ -515,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 }) @@ -540,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 { @@ -569,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)) @@ -585,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 { @@ -647,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")) @@ -695,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})) } }) }) @@ -718,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)) @@ -738,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...) } @@ -758,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") @@ -769,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})) } }) @@ -777,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 { @@ -794,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...) @@ -809,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 @@ -829,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()) @@ -849,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)) }) @@ -862,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()) }) @@ -885,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()) }) @@ -912,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() { @@ -919,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)) }) @@ -975,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)) @@ -1022,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)) }) @@ -1046,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 { @@ -1069,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)) @@ -1077,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()) }) }) }) @@ -1093,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 @@ -1109,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()) @@ -1127,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()) }) @@ -1140,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)) @@ -1172,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)) @@ -1213,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()) @@ -1230,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.go b/persistence/criteria_sql.go index 890f757d7..d2f817438 100644 --- a/persistence/criteria_sql.go +++ b/persistence/criteria_sql.go @@ -13,6 +13,7 @@ import ( "github.com/Masterminds/squirrel" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/criteria" + "golang.org/x/text/unicode/norm" ) type smartPlaylistJoinType int @@ -344,11 +345,15 @@ func startOfPeriod(numDays int64, from time.Time) string { } func (c smartPlaylistCriteria) inList(values map[string]any, negate bool) (squirrel.Sqlizer, error) { - playlistID, ok := values["id"].(string) - if !ok { - return nil, errors.New("playlist id not given") + var condition squirrel.Sqlizer + if playlistId, ok := values["id"].(string); ok && playlistId != "" { + condition = squirrel.Eq{"pl.playlist_id": playlistId} + } else if playlistPath, ok := values["path"].(string); ok && playlistPath != "" { + condition = squirrel.Eq{"playlist.path": pathVariants(playlistPath)} + } else { + return nil, errors.New("playlist id or path not given") } - filters := squirrel.And{squirrel.Eq{"pl.playlist_id": playlistID}} + filters := squirrel.And{condition} if !c.owner.IsAdmin { if c.owner.ID == "" { filters = append(filters, squirrel.Eq{"playlist.public": 1}) @@ -373,6 +378,18 @@ func (c smartPlaylistCriteria) inList(values map[string]any, negate bool) (squir return squirrel.Expr("media_file.id IN ("+subSQL+")", subArgs...), nil } +// Filesystems disagree on the Unicode form of a name, so match the path in NFC and NFD. +func pathVariants(path string) []string { + variants := []string{path} + if alt := norm.NFC.String(path); alt != path { + variants = append(variants, alt) + } + if alt := norm.NFD.String(path); alt != path { + variants = append(variants, alt) + } + return variants +} + func jsonExpr(info criteria.FieldInfo, cond squirrel.Sqlizer, negate bool) squirrel.Sqlizer { if info.IsRole { return roleCond{role: info.Name(), cond: cond, not: negate} 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/criteria_sql_test.go b/persistence/criteria_sql_test.go index 8d0069b0d..ef91ec989 100644 --- a/persistence/criteria_sql_test.go +++ b/persistence/criteria_sql_test.go @@ -47,7 +47,10 @@ var _ = Describe("Smart playlist criteria SQL", func() { Entry("in range", criteria.InTheRange{"year": []int{1980, 1990}}, "(media_file.year >= ? AND media_file.year <= ?)", 1980, 1990), Entry("before", criteria.Before{"lastPlayed": time.Date(2021, 10, 1, 0, 0, 0, 0, time.Local)}, "annotation.play_date < ?", time.Date(2021, 10, 1, 0, 0, 0, 0, time.Local)), Entry("after", criteria.After{"lastPlayed": time.Date(2021, 10, 1, 0, 0, 0, 0, time.Local)}, "annotation.play_date > ?", time.Date(2021, 10, 1, 0, 0, 0, 0, time.Local)), - Entry("in playlist", criteria.InPlaylist{"id": "deadbeef-dead-beef"}, "media_file.id IN (SELECT media_file_id FROM playlist_tracks pl LEFT JOIN playlist on pl.playlist_id = playlist.id WHERE (pl.playlist_id = ? AND playlist.public = ?))", "deadbeef-dead-beef", 1), + Entry("in playlist [path]", criteria.InPlaylist{"path": "lacuslacus.nsp"}, "media_file.id IN (SELECT media_file_id FROM playlist_tracks pl LEFT JOIN playlist on pl.playlist_id = playlist.id WHERE (playlist.path IN (?) AND playlist.public = ?))", "lacuslacus.nsp", 1), + Entry("in playlist [id]", criteria.InPlaylist{"id": "deadbeef-dead-beef"}, "media_file.id IN (SELECT media_file_id FROM playlist_tracks pl LEFT JOIN playlist on pl.playlist_id = playlist.id WHERE (pl.playlist_id = ? AND playlist.public = ?))", "deadbeef-dead-beef", 1), + Entry("in playlist [empty id falls back to path]", criteria.InPlaylist{"id": "", "path": "/music/x.nsp"}, "media_file.id IN (SELECT media_file_id FROM playlist_tracks pl LEFT JOIN playlist on pl.playlist_id = playlist.id WHERE (playlist.path IN (?) AND playlist.public = ?))", "/music/x.nsp", 1), + Entry("in playlist [decomposed unicode path]", criteria.InPlaylist{"path": "/m\u00fasica/x.nsp"}, "media_file.id IN (SELECT media_file_id FROM playlist_tracks pl LEFT JOIN playlist on pl.playlist_id = playlist.id WHERE (playlist.path IN (?,?) AND playlist.public = ?))", "/m\u00fasica/x.nsp", "/mu\u0301sica/x.nsp", 1), Entry("not in playlist", criteria.NotInPlaylist{"id": "deadbeef-dead-beef"}, "media_file.id NOT IN (SELECT media_file_id FROM playlist_tracks pl LEFT JOIN playlist on pl.playlist_id = playlist.id WHERE (pl.playlist_id = ? AND playlist.public = ?))", "deadbeef-dead-beef", 1), Entry("album annotation", criteria.Gt{"albumRating": 3}, "album_annotation.rating > ?", 3), Entry("artist annotation", criteria.Is{"artistLoved": true}, "artist_annotation.starred = ?", true), @@ -268,6 +271,13 @@ var _ = Describe("Smart playlist criteria SQL", func() { Expect(err).To(MatchError(ContainSubstring("invalid boolean value for 'missing' expression"))) }) + It("returns an error when inPlaylist has empty path", func() { + _, err := newSmartPlaylistCriteria( + criteria.Criteria{Expression: criteria.InPlaylist{"path": ""}}, + withSmartPlaylistOwner(model.User{ID: "owner-id", IsAdmin: false})).where() + Expect(err).To(MatchError(ContainSubstring("playlist id or path not given"))) + }) + It("returns an error for a range over a tag/role field", func() { _, err := newSmartPlaylistCriteria(criteria.Criteria{Expression: criteria.InTheRange{"rate": []int{1, 5}}}).where() Expect(err).To(MatchError(ContainSubstring("range operator not supported for tag/role field"))) 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 a4b73d9d6..b4f7c9069 100644 --- a/persistence/folder_repository.go +++ b/persistence/folder_repository.go @@ -27,6 +27,14 @@ type dbFolder struct { ImageFiles string `structs:"-" json:"-"` } +// String guards the promoted Folder.String(), which would dereference a nil Folder. +func (f dbFolder) String() string { + if f.Folder == nil { + return "" + } + return f.Folder.String() +} + func (f *dbFolder) PostScan() error { var err error if f.ImageFiles != "" { @@ -56,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) } } @@ -114,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 } @@ -125,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}, @@ -164,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 @@ -177,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 } @@ -222,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}, @@ -245,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 } @@ -266,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 } @@ -295,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}, @@ -307,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 545d8ad49..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,58 @@ 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()) + }) + }) + + Describe("dbFolder.String", func() { + It("does not dereference a nil Folder", func() { + Expect(fmt.Sprint(dbFolder{})).To(Equal("")) + Expect(fmt.Sprint(&dbFolder{})).To(Equal("")) }) }) @@ -360,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. @@ -371,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..85da65cbc 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 { @@ -94,10 +93,12 @@ func (r *libraryRepository) Put(l *model.Library, colsToUpdate ...string) error "path": l.Path, "remote_path": l.RemotePath, "default_new_users": l.DefaultNewUsers, + "pid_album": l.PIDAlbum, + "pid_track": l.PIDTrack, }, 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 +123,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 +135,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 +149,86 @@ 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) SetScannedPID(ctx context.Context, id int, pid model.PIDConfig) error { + sq := Update(r.tableName). + Set("scanned_pid_album", pid.Album). + Set("scanned_pid_track", pid.Track). + Where(Eq{"id": id}) + _, err := r.executeSQL(ctx, sq) + return err +} + +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 +246,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 +275,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 +302,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..bf485a06f 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,22 +257,55 @@ 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)) }) }) }) + Describe("PID config", func() { + It("stores the overrides, and Put never touches the scanned specs", func() { + lib := &model.Library{Name: "PID Library", Path: "/music/pid", PIDAlbum: "folder", PIDTrack: "title"} + Expect(repo.Put(ctx, lib)).To(Succeed()) + Expect(repo.SetScannedPID(ctx, lib.ID, model.PIDConfig{Album: "folder", Track: "title"})).To(Succeed()) + + // An update coming from the REST API has no scanned specs. It must not clear them + update := &model.Library{ID: lib.ID, Name: "PID Library", Path: "/music/pid", PIDTrack: "title"} + Expect(repo.Put(ctx, update)).To(Succeed()) + + saved, err := repo.Get(ctx, lib.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(saved.PIDAlbum).To(BeEmpty()) + Expect(saved.PIDTrack).To(Equal("title")) + Expect(saved.ScannedPIDAlbum).To(Equal("folder")) + Expect(saved.ScannedPIDTrack).To(Equal("title")) + }) + + It("keeps the overrides when a partial update does not send them", func() { + lib := &model.Library{Name: "Partial", Path: "/music/partial", PIDAlbum: "folder", PIDTrack: "title"} + Expect(repo.Put(ctx, lib)).To(Succeed()) + + Expect(repo.Put(ctx, &model.Library{ID: lib.ID, Name: "Renamed"}, "name")).To(Succeed()) + + saved, err := repo.Get(ctx, lib.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(saved.Name).To(Equal("Renamed")) + Expect(saved.PIDAlbum).To(Equal("folder")) + Expect(saved.PIDTrack).To(Equal("title")) + }) + }) + 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 +316,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 da167b1a8..b9a5c3374 100644 --- a/persistence/mediafile_repository.go +++ b/persistence/mediafile_repository.go @@ -38,6 +38,14 @@ type dbMediaFile struct { RgTrackPeak *float64 `structs:"-" json:"-"` } +// String guards the promoted MediaFile.String(), which would dereference a nil MediaFile. +func (m dbMediaFile) String() string { + if m.MediaFile == nil { + return "" + } + return m.MediaFile.String() +} + func (m *dbMediaFile) PostScan() error { m.RGTrackGain = m.RgTrackGain m.RGTrackPeak = m.RgTrackPeak @@ -77,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()) @@ -139,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 } @@ -168,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 } @@ -213,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] @@ -248,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 { @@ -264,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 { @@ -296,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 } @@ -309,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. @@ -331,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...) @@ -340,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 } @@ -355,26 +362,21 @@ 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{} var unqualified []string for _, path := range paths { - parts := strings.SplitN(path, ":", 2) - if len(parts) == 2 { - // Library-qualified path: "libraryID:path" - libraryID, err := strconv.Atoi(parts[0]) - if err != nil { - // Invalid format, skip - continue + // A numeric prefix is ambiguous: "1:foo.mp3" qualifies a library, but "1999: A Life/01.mp3" + // is a plain path. Search both ways rather than guessing. + if id, rest, ok := strings.Cut(path, ":"); ok { + if libraryID, err := strconv.Atoi(id); err == nil { + byLibrary[libraryID] = append(byLibrary[libraryID], rest) } - byLibrary[libraryID] = append(byLibrary[libraryID], parts[1]) - } else { - // Unqualified path: search across all libraries - unqualified = append(unqualified, path) } + unqualified = append(unqualified, path) } query := Or{} @@ -392,34 +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) DeleteAllMissing() (int64, error) { - user := loggedUser(r.ctx) +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(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(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(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(ctx, buf); err != nil { + return fmt.Errorf("reassigning buffered scrobbles: %w", err) + } + return nil +} + +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}, @@ -427,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). @@ -453,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 } @@ -466,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}, @@ -476,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{ @@ -484,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 } @@ -497,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}, @@ -510,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 } @@ -519,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}, @@ -536,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 } @@ -549,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 { @@ -560,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 8a492a813..14590090c 100644 --- a/persistence/mediafile_repository_test.go +++ b/persistence/mediafile_repository_test.go @@ -16,6 +16,7 @@ import ( "github.com/navidrome/navidrome/model/criteria" "github.com/navidrome/navidrome/model/id" "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/utils/slice" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" "github.com/pocketbase/dbx" @@ -23,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() { @@ -36,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() @@ -59,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")) }) @@ -81,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()) @@ -128,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()) @@ -137,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))) }) }) @@ -150,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 @@ -181,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{ @@ -205,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 @@ -228,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"}, }) @@ -247,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)) @@ -258,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 { @@ -275,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)) @@ -284,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 { @@ -303,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() { @@ -320,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() { @@ -339,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{ @@ -349,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())) @@ -424,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))) @@ -510,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())) @@ -554,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 @@ -596,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) } }) @@ -606,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}}, @@ -629,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}}, @@ -650,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}}, @@ -676,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) @@ -686,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()) @@ -708,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 { @@ -734,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)) @@ -748,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 { @@ -765,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()) }) }) @@ -778,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 { @@ -787,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 { @@ -796,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()) }) @@ -819,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")) @@ -837,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")) @@ -845,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()) }) @@ -861,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 { @@ -885,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...) } @@ -911,13 +909,113 @@ 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()) }) }) }) + Describe("ReassignReferences", func() { + var prev, next model.MediaFile + var pr model.PlaylistRepository + var pls model.Playlist + + BeforeEach(func() { + ctx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) + 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(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(ctx, &pls)).To(Succeed()) + }) + + AfterEach(func() { + _ = 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(ctx, 5, prev.ID)).To(Succeed()) + Expect(mr.AddBookmark(ctx, prev.ID, "here", 42)).To(Succeed()) + + Expect(mr.ReassignReferences(ctx, prev.ID, next.ID)).To(Succeed()) + + got, err := mr.Get(ctx, next.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(got.Rating).To(Equal(5)) + + bookmarks, err := mr.GetBookmarks(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(bookmarks).To(ContainElement(HaveField("Item.ID", next.ID))) + + 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)) + }) + + It("moves scrobbles and buffered scrobbles onto the new id", func() { + ctx := request.WithUser(log.NewContext(context.TODO()), model.User{ID: "userid"}) + 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(ctx, prev.ID, next.ID)).To(Succeed()) + Expect(mr.Delete(ctx, prev.ID)).To(Succeed()) + + 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(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() { + 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(ctx, prev.ID, next.ID)).To(Succeed()) + + 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(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(ctx, prev.ID, next.ID)).To(Succeed()) + + got, err := mr.Get(ctx, next.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(got.Rating).To(Equal(1)) + bookmarks, err := mr.GetBookmarks(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(bookmarks).To(ContainElement(SatisfyAll(HaveField("Item.ID", next.ID), HaveField("Comment", "next")))) + }) + }) + Describe("FindByPaths", func() { // Test fixtures for Unicode and case-sensitivity tests var testFiles []model.MediaFile @@ -930,20 +1028,43 @@ var _ = Describe("MediaRepository", func() { {ID: "findpath-3", LibraryID: 1, Path: "plex/02 - ACROSS.flac", Title: "Fullwidth"}, // French diacritic: è (U+00E8, can decompose to e + combining grave) {ID: "findpath-4", LibraryID: 1, Path: "artist/Michèle/song.mp3", Title: "French"}, + {ID: "findpath-5", LibraryID: 1, Path: "Bach: Goldberg Variations/01.mp3", Title: "Colon"}, + {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(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(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(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")) @@ -951,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")) @@ -960,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", }) @@ -982,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()) }) @@ -1011,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), }) @@ -1068,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", }) @@ -1080,6 +1205,13 @@ var _ = Describe("MediaRepository", func() { }) }) + Describe("dbMediaFile.String", func() { + It("does not dereference a nil MediaFile", func() { + Expect(fmt.Sprint(dbMediaFile{})).To(Equal("")) + Expect(fmt.Sprint(&dbMediaFile{})).To(Equal("")) + }) + }) + Describe("wrapMediaFileCursor", func() { It("does not panic when the cursor yields a dbMediaFile with nil MediaFile", func() { // Simulate what queryWithStableResults does on the rows.Err() path: @@ -1124,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()) @@ -1144,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() { @@ -1152,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) }) }) @@ -1196,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 9d3a33cfc..44e944bff 100644 --- a/persistence/persistence.go +++ b/persistence/persistence.go @@ -3,7 +3,8 @@ package persistence import ( "context" "database/sql" - "reflect" + "fmt" + "sync" "time" "github.com/navidrome/navidrome/db" @@ -14,135 +15,155 @@ 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) ArtworkQueue() model.ArtworkQueueRepository { + return s.artworkQueue() } -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) +func scopeLabel(scope []string) string { + if len(scope) > 0 { + return scope[0] } - log.Error("Resource not implemented", "model", reflect.TypeOf(m).Name()) - return nil + return "" } func (s *SQLStore) WithTx(block func(tx model.DataStore) error, scope ...string) error { - var msg string - if len(scope) > 0 { - msg = scope[0] - } + msg := scopeLabel(scope) start := time.Now() conn, inTx := s.db.(*dbx.DB) if !inTx { @@ -152,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) @@ -168,15 +189,60 @@ 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) }, 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 { @@ -194,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 { - log.Error(ctx, "Error tidying up database", err) + return fmt.Errorf("tidying up database: %w", err) } - return err -} - -func (s *SQLStore) getDBXBuilder() dbx.Builder { - if s.db == nil { - return dbx.NewFromDB(db.Db(), db.Driver) - } - return s.db + return nil } diff --git a/persistence/persistence_suite_test.go b/persistence/persistence_suite_test.go index f146cb06b..ee2794454 100644 --- a/persistence/persistence_suite_test.go +++ b/persistence/persistence_suite_test.go @@ -169,14 +169,36 @@ func p(path string) string { return filepath.FromSlash(path) } +// restrictedFixture creates a second library plus a non-admin user granted library 1 only, so +// specs can assert that a query filters by library. Cleans itself up after the spec. +func restrictedFixture(name string) (context.Context, model.Library, model.User) { + adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) + db := GetDBXBuilder() + + lib := model.Library{Name: name + " Library", Path: "/" + name} + lr := NewLibraryRepository(db) + Expect(lr.Put(adminCtx, &lib)).To(Succeed()) + + user := createUserWithLibraries(name+"-restricted", []int{1}) + ur := NewUserRepository(db) + Expect(ur.Put(adminCtx, &user)).To(Succeed()) + Expect(ur.SetUserLibraries(adminCtx, user.ID, []int{1})).To(Succeed()) + + DeferCleanup(func() { + _ = NewUserRepository(db).Delete(adminCtx, user.ID) + _ = NewLibraryRepository(db).(*libraryRepository).delete(adminCtx, squirrel.Eq{"id": lib.ID}) + }) + return adminCtx, lib, user +} + var _ = BeforeSuite(func() { conn := GetDBXBuilder() 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) } @@ -184,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) } @@ -225,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", @@ -236,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) } @@ -265,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) } @@ -288,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) } @@ -302,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) } @@ -313,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 13e56bde1..43d2e81ed 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" @@ -20,39 +23,109 @@ 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)) }) }) }) + + 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().Put(ctx, "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().Get(ctx, "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().Put(ctx, "outer-key", "v")).To(Succeed()) + Expect(tx.WithTxRetry(ctx, func(ctx context.Context, inner model.DataStore) error { + 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().Get(ctx, "inner-key") + Expect(err).To(MatchError(model.ErrNotFound)) + }) + }) }) diff --git a/persistence/player_repository.go b/persistence/player_repository.go index 353b0444f..ba5d26794 100644 --- a/persistence/player_repository.go +++ b/persistence/player_repository.go @@ -2,11 +2,16 @@ package persistence import ( "context" - "errors" + "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" ) @@ -14,12 +19,12 @@ 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"), + "name": containsFilter("player.name"), + "hasapikey": hasAPIKeyFilter, }) r.setSortMappings(map[string]string{ "user_name": "username", //TODO rename all user_name and userName to username @@ -27,43 +32,50 @@ 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 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 } -func (r *playerRepository) selectPlayer(options ...model.QueryOptions) SelectBuilder { - return r.newSelect(options...). - Columns("player.*"). +func (r *playerRepository) selectPlayer(ctx context.Context, options ...model.QueryOptions) SelectBuilder { + return r.newSelect(ctx, options...). + 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") } -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", @@ -72,7 +84,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 } @@ -83,67 +95,137 @@ 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" +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 (r *playerRepository) NewInstance() any { - return &model.Player{} +func validateAPIKey(key string) error { + if !apiKeyFormat.MatchString(key) { + return apiKeyValidationError("resources.player.validation.apiKeyFormat") + } + return nil } -// 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) - 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) { + u := loggedUser(ctx) + if t.UserId == "" && u.ID != invalidUserId { + t.UserId = u.ID + } + if t.UserId != u.ID { return "", rest.ErrPermissionDenied } - id, err := r.put(t.ID, t) - if errors.Is(err, model.ErrNotFound) { - return "", rest.ErrNotFound + // Hand-made players are only reachable through a key, so one is required + if t.APIKey == nil || *t.APIKey == "" { + return "", apiKeyValidationError("ra.validation.required") } - return id, err + 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(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...) + 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) Delete(id string) error { - return r.deleteOwned(id) +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 = (*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 b7085a1fb..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,9 +13,18 @@ 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 + 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 +35,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 +105,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 +113,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 +169,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,11 +187,12 @@ var _ = Describe("PlayerRepository", func() { clone := player clone.ID = "" clone.IP = "192.168.1.1" - id, err := repo.Save(&clone) + 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 { @@ -197,11 +200,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)) @@ -209,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{}), ) }) @@ -223,7 +227,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 +242,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)) @@ -258,13 +262,266 @@ var _ = Describe("PlayerRepository", func() { Entry("regular context", false, model.Players{regularPlayer}, regularPlayer, adminPlayer1), ) - Describe("Ownership enforcement (cross-tenant write protection)", func() { - var regularRepo *playerRepository + Describe("API keys", func() { + const key = testAPIKey + const otherKey = "nds_ABCDEFGHIJKLMNOPQRSTUV" + var ownerCtx, otherCtx context.Context BeforeEach(func() { - ctx := log.NewContext(context.TODO()) - ctx = request.WithUser(ctx, regularUser) - regularRepo = NewPlayerRepository(ctx, database).(*playerRepository) + 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 + + BeforeEach(func() { + 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,15 +536,37 @@ 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)) }) + It("does not let a regular user overwrite another user's player via Save with a spoofed id", func() { + spoofed := model.Player{ + ID: adminPlayer1.ID, + Name: "HIJACKED", + UserId: regularUser.ID, + ReportRealPath: true, + APIKey: new(testAPIKey), + } + + id, err := regularRepo.Save(regularCtx, &spoofed) + Expect(err).To(BeNil()) + Expect(id).ToNot(Equal(adminPlayer1.ID)) + + stored, err := adminRepo.Get(ctx, adminPlayer1.ID) + Expect(err).To(BeNil()) + Expect(*stored).To(Equal(adminPlayer1)) + + created, err := adminRepo.Get(ctx, id) + Expect(err).To(BeNil()) + Expect(created.UserId).To(Equal(regularUser.ID)) + }) + It("does not let a regular user reassign their own player to another user", func() { // Owner updates their own player but tries to give it away to the admin. The update // succeeds for the other fields, but user_id is never written, so ownership stays put. @@ -295,11 +574,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)) }) @@ -310,11 +589,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)) @@ -324,10 +603,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)) @@ -335,7 +614,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 bc6bffd25..41b75266c 100644 --- a/persistence/playlist_repository.go +++ b/persistence/playlist_repository.go @@ -13,6 +13,7 @@ import ( "github.com/deluan/rest" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/utils/slice" "github.com/pocketbase/dbx" ) @@ -49,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"), @@ -80,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{} } @@ -91,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 == "" @@ -122,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 } @@ -134,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 } @@ -184,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 } @@ -204,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 @@ -250,36 +250,43 @@ 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 } - return r.addTracks(playlistId, 1, mediaFileIds) + _, err = r.addTracks(ctx, playlistId, 1, mediaFileIds) + return err } -func (r *playlistRepository) addTracks(playlistId string, startingPos int, mediaFileIds []string) error { +// 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(ctx context.Context, playlistId string, startingPos int, mediaFileIds []string) (int, error) { + mediaFileIds, err := r.keepAccessible(ctx, mediaFileIds) + if err != nil { + return 0, err + } // Break the track list in chunks to avoid hitting SQLITE_MAX_VARIABLE_NUMBER limit // Add new tracks, chunk by chunk pos := startingPos @@ -289,18 +296,40 @@ func (r *playlistRepository) addTracks(playlistId string, startingPos int, media ins = ins.Values(playlistId, t, pos) pos++ } - _, err := r.executeSQL(ins) - if err != nil { - return err + if _, err := r.executeSQL(ctx, ins); err != nil { + return 0, err } } - r.enqueueCoverRebuild(playlistId) - return 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(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(ctx, Select("id").From("media_file").Where(Eq{"id": chunk}), "media_file") + var found []string + if err := r.queryAllSlice(ctx, sq, &found); err != nil { + return nil, err + } + for _, id := range found { + accessible[id] = struct{}{} + } + } + return slice.Filter(mediaFileIds, func(id string) bool { + _, ok := accessible[id] + return ok + }), nil } // 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", @@ -310,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 } @@ -323,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 } @@ -336,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", @@ -370,59 +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")...) - if errors.Is(err, model.ErrNotFound) { - return rest.ErrNotFound - } + _, 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"). @@ -430,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) } } @@ -458,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 @@ -468,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 = ? @@ -479,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 cf1b8f3fa..392446cef 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,36 +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}) - 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'"+ @@ -113,22 +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(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 } @@ -136,85 +138,108 @@ 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(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 } - err := r.queryOne(sq, &res) - if err != nil { + if err := r.queryOne(ctx, sq, &res); err != nil { return 0, err } - return len(mediaFileIds), r.playlistRepo.addTracks(r.playlistId, int(res.Max.Int32+1), mediaFileIds) + return r.playlistRepo.addTracks(ctx, r.playlistId, int(res.Max.Int32+1), mediaFileIds) } -func (r *playlistTrackRepository) addMediaFileIds(cond Sqlizer) (int, error) { +// 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(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(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(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(ctx, 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(ctx, r.playlistId) +} + +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) { + // 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(discs []model.DiscID) (int, error) { +func (r *playlistTrackRepository) AddDiscs(ctx context.Context, discs []model.DiscID) (int, error) { if len(discs) == 0 { return 0, nil } @@ -222,40 +247,50 @@ 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. -func (r *playlistTrackRepository) Reorder(pos int, newPos int) error { +// 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(ctx context.Context, pos int, newPos int) error { + var res struct{ Max sql.NullInt32 } + 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) + if pos < 1 || pos > last { + return model.ErrNotFound + } + newPos = min(max(newPos, 1), last) if pos == newPos { return nil } 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 @@ -263,11 +298,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)) } @@ -276,14 +311,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 88ddbff48..119c57116 100644 --- a/persistence/playlist_track_repository_test.go +++ b/persistence/playlist_track_repository_test.go @@ -1,11 +1,13 @@ package persistence import ( + "context" "strconv" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/utils/slice" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -15,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))) }) }) @@ -46,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)) }) @@ -58,26 +61,121 @@ 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})) }) }) + Describe("Insert", func() { + var tracks model.PlaylistTrackRepository + + BeforeEach(func() { + plsRepo := NewPlaylistRepository(GetDBXBuilder()) + pls := model.Playlist{Name: "Insert", OwnerID: "userid", OwnerName: "userid"} + Expect(plsRepo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { Expect(plsRepo.Delete(ctx, pls.ID)).To(Succeed()) }) + + 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(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(ctx, []string{songComeTogether.ID, songAntenna.ID}, pos)).To(Equal(2)) + Expect(order()).To(Equal(want())) + 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} + }), + Entry("at the start, for zero or less", 0, func() []string { + return []string{songComeTogether.ID, songAntenna.ID, songDayInALife.ID, songRadioactivity.ID} + }), + Entry("at the end, past the last position", 9, func() []string { + return []string{songDayInALife.ID, songRadioactivity.ID, songComeTogether.ID, songAntenna.ID} + }), + ) + + It("renumbers positions contiguously", func() { + 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"})) + }) + }) + + Describe("Reorder", func() { + var tracks model.PlaylistTrackRepository + + BeforeEach(func() { + plsRepo := NewPlaylistRepository(GetDBXBuilder()) + pls := model.Playlist{Name: "Reorder", OwnerID: "userid", OwnerName: "userid"} + Expect(plsRepo.Put(ctx, &pls)).To(Succeed()) + DeferCleanup(func() { Expect(plsRepo.Delete(ctx, pls.ID)).To(Succeed()) }) + + 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(ctx, model.QueryOptions{Sort: "id"}) + Expect(err).ToNot(HaveOccurred()) + var ids, songs []string + for _, t := range all { + ids = append(ids, t.ID) + songs = append(songs, t.MediaFileID) + } + return ids, songs + } + + DescribeTable("clamps the destination to the playlist", + func(newPos int, want func() []string) { + Expect(tracks.Reorder(ctx, 1, newPos)).To(Succeed()) + ids, songs := rows() + Expect(ids).To(Equal([]string{"1", "2", "3"})) + Expect(songs).To(Equal(want())) + }, + Entry("past the end moves to the end", 9, func() []string { + return []string{songRadioactivity.ID, songComeTogether.ID, songDayInALife.ID} + }), + Entry("below 1 stays first", -4, func() []string { + return []string{songDayInALife.ID, songRadioactivity.ID, songComeTogether.ID} + }), + ) + + DescribeTable("rejects a source position outside the playlist, leaving rows untouched", + func(pos int) { + 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})) + }, + Entry("past the end", 4), + Entry("zero", 0), + ) + }) + Describe("Delete", func() { var tracks model.PlaylistTrackRepository const numTracks = deleteChunkSize*2 + 1 @@ -91,35 +189,219 @@ 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()) + }) + }) + + 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 + var adminCtx, userCtx context.Context + var userTracks model.PlaylistTrackRepository + var plsID string + + BeforeEach(func() { + adminCtx, otherLib, restrictedUser = restrictedFixture("pls") + userCtx = request.WithUser(log.NewContext(GinkgoT().Context()), restrictedUser) + db := GetDBXBuilder() + + 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(adminCtx, "pls-otherlib-track") }) + + adminPls := NewPlaylistRepository(db) + pls := model.Playlist{Name: "Public Mixed", OwnerID: adminUser.ID, OwnerName: adminUser.UserName, Public: true} + Expect(adminPls.Put(adminCtx, &pls)).To(Succeed()) + plsID = pls.ID + 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(db).Tracks(userCtx, plsID, false) + }) + + It("Read does not return a track outside the user's libraries", func() { + _, 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(userCtx, "1") + Expect(err).ToNot(HaveOccurred()) + Expect(trk.MediaFile.ID).To(Equal(songDayInALife.ID)) + }) + + It("Count excludes tracks outside the user's libraries", func() { + 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(userCtx)).ToNot(ContainElement("pls-hidden-album")) + }) + + Describe("Add", func() { + var ownTracks model.PlaylistTrackRepository + + BeforeEach(func() { + userPls := NewPlaylistRepository(GetDBXBuilder()) + own := model.Playlist{Name: "Own Playlist", OwnerID: restrictedUser.ID, OwnerName: restrictedUser.UserName} + 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(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(userCtx, []string{"pls-hidden-album"})).To(BeZero()) + }) + + It("drops them when reached through Insert", func() { + Expect(ownTracks.Add(userCtx, []string{songDayInALife.ID})).To(Equal(1)) + + Expect(ownTracks.Insert(userCtx, []string{"pls-otherlib-track", songComeTogether.ID}, 1)).To(Equal(1)) + + 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") + }) + }) + + Describe("Put", func() { + storedIDs := func(id string) []string { + ids, err := NewPlaylistRepository(GetDBXBuilder()).Tracks(adminCtx, id, false).GetMediaFileIDs(adminCtx) + Expect(err).ToNot(HaveOccurred()) + return ids + } + put := func(ctx context.Context, owner model.User, pls *model.Playlist, ids ...string) string { + pls.OwnerID = owner.ID + pls.Tracks = nil + pls.AddMediaFilesByID(ids) + Expect(NewPlaylistRepository(GetDBXBuilder()).Put(ctx, pls)).To(Succeed()) + DeferCleanup(func() { _ = NewPlaylistRepository(GetDBXBuilder()).Delete(adminCtx, pls.ID) }) + return pls.ID + } + + It("drops ids outside the user's libraries when creating a playlist", func() { + id := put(userCtx, restrictedUser, &model.Playlist{Name: "Created"}, songDayInALife.ID, "pls-otherlib-track") + + Expect(storedIDs(id)).To(Equal([]string{songDayInALife.ID})) + }) + + It("drops them when replacing the tracks of an existing playlist", func() { + pls := &model.Playlist{Name: "Replaced"} + put(userCtx, restrictedUser, pls, songDayInALife.ID) + + put(userCtx, restrictedUser, pls, "pls-otherlib-track") + + Expect(storedIDs(pls.ID)).To(BeEmpty()) + }) + + It("does not count a dropped id, so it cannot be told apart from an unknown one", 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(GetDBXBuilder()) + h, err := userPls.Get(userCtx, hidden) + Expect(err).ToNot(HaveOccurred()) + u, err := userPls.Get(userCtx, unknown) + Expect(err).ToNot(HaveOccurred()) + Expect(h.SongCount).To(Equal(u.SongCount)) + Expect(h.Duration).To(Equal(u.Duration)) + Expect(h.Size).To(Equal(u.Size)) + }) + + It("keeps order and duplicates of the accessible ids", func() { + id := put(userCtx, restrictedUser, &model.Playlist{Name: "Ordered"}, + songDayInALife.ID, "pls-otherlib-track", songComeTogether.ID, songDayInALife.ID) + + Expect(storedIDs(id)).To(Equal([]string{songDayInALife.ID, songComeTogether.ID, songDayInALife.ID})) + }) + + It("keeps every id when run as an admin, as the scanner's playlist sync does", func() { + id := put(adminCtx, adminUser, &model.Playlist{Name: "Synced"}, songDayInALife.ID, "pls-otherlib-track") + + Expect(storedIDs(id)).To(Equal([]string{songDayInALife.ID, "pls-otherlib-track"})) + }) + }) + + It("still shows everything to an admin", func() { + 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 c1e36f0b1..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,56 +25,64 @@ 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 + 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 } -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 } @@ -121,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 dc68b0892..44330250b 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" @@ -10,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)) }) }) @@ -61,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)) @@ -71,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)) }) @@ -112,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)) @@ -133,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"}`)) @@ -156,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")) }) @@ -170,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")) }) @@ -178,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) @@ -193,52 +200,63 @@ 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(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) + _, err = conn.ExecContext(GinkgoT().Context(), "BEGIN IMMEDIATE") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _, _ = conn.ExecContext(context.Background(), "ROLLBACK") }) + + 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 e042ee6bf..b5d4f3a07 100644 --- a/persistence/radio_repository.go +++ b/persistence/radio_repository.go @@ -2,7 +2,6 @@ package persistence import ( "context" - "errors" "time" . "github.com/Masterminds/squirrel" @@ -17,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"), @@ -27,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.delete(Eq{"id": 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 } @@ -92,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 } @@ -100,57 +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) - if errors.Is(err, model.ErrNotFound) { - return "", rest.ErrNotFound - } + 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 } - err := r.Put(t, cols...) - if errors.Is(err, model.ErrNotFound) { - return rest.ErrNotFound - } - return err + 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 a958f715d..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,31 +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(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)) }) @@ -71,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)) @@ -80,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", @@ -88,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"), @@ -134,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")) @@ -151,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)) }) @@ -172,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)) }) @@ -187,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)) @@ -196,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 4b6dc9240..6dd9c3d85 100644 --- a/persistence/share_repository.go +++ b/persistence/share_repository.go @@ -2,7 +2,6 @@ package persistence import ( "context" - "errors" "fmt" "strings" "time" @@ -19,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{ @@ -30,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) } @@ -71,8 +69,7 @@ func (r *shareRepository) GetAll(options ...model.QueryOptions) (model.Shares, e return res, err } -func (r *shareRepository) loadMedia(share *model.Share) error { - var err error +func (r *shareRepository) loadMedia(ctx context.Context, share *model.Share) error { ids := strings.Split(share.ResourceIDs, ",") if len(ids) == 0 { return nil @@ -80,73 +77,68 @@ func (r *shareRepository) loadMedia(share *model.Share) error { noMissing := func(cond Sqlizer) Sqlizer { return And{cond, Eq{"missing": false}} } + // Load as the share owner so their library access is applied, whoever renders the share. + ownerCtx, err := r.ownerContext(ctx, share) + if err != nil { + return err + } switch share.ResourceType { 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. - // Load as the share owner so their library access is applied. - ctx, err := r.ownerContext(share) + 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 } - albumRepo := NewAlbumRepository(ctx, r.db) - share.Albums, err = albumRepo.GetAll(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(r.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(r.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": - // Load tracks as the share owner so their library access is applied. - ctx, err := r.ownerContext(share) - if err != nil { - return err - } - 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(r.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 { @@ -163,62 +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 - u := loggedUser(r.ctx) - if s.UserID == "" { + // 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(ctx) + if u.ID != invalidUserId || s.UserID == "" { s.UserID = u.ID } s.CreatedAt = time.Now() s.UpdatedAt = time.Now() - id, err := r.put(s.ID, s) - if errors.Is(err, model.ErrNotFound) { - return "", rest.ErrNotFound - } - return id, err + 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 3af91b2af..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,14 +228,14 @@ 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()) }) }) - Describe("Artist share library scoping", func() { + Describe("Artist, album and media file share library scoping", func() { var otherLib model.Library var owner model.User const primaryID = "share-aa-primary" @@ -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,50 +260,57 @@ 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()) - _, err := b.NewQuery(` - INSERT INTO share (id, user_id, description, resource_type, resource_ids, created_at, updated_at) - VALUES ({:id}, {:user}, {:desc}, {:type}, {:ids}, {:created}, {:updated}) - `).Bind(map[string]any{ - "id": "art-share", "user": owner.ID, "desc": "Artist scope share", - "type": "artist", "ids": secondaryID, "created": time.Now(), "updated": time.Now(), - }).Execute() - Expect(err).ToNot(HaveOccurred()) + for _, s := range []struct{ id, typ, ids string }{ + {"art-share", "artist", secondaryID}, + {"art-album-share", "album", "art-album-ok,art-album-other"}, + {"art-mf-share", "media_file", "art-ok,art-other"}, + } { + _, err := b.NewQuery(` + INSERT INTO share (id, user_id, description, resource_type, resource_ids, created_at, updated_at) + VALUES ({:id}, {:user}, {:desc}, {:type}, {:ids}, {:created}, {:updated}) + `).Bind(map[string]any{ + "id": s.id, "user": owner.ID, "desc": "Scope share", + "type": s.typ, "ids": s.ids, "created": time.Now(), "updated": time.Now(), + }).Execute() + Expect(err).ToNot(HaveOccurred()) + } }) AfterEach(func() { adminCtx := request.WithUser(log.NewContext(GinkgoT().Context()), adminUser) b := GetDBXBuilder() - _, _ = b.NewQuery(`DELETE FROM share WHERE id = 'art-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) + _, _ = b.NewQuery(`DELETE FROM share WHERE id IN ('art-share', 'art-album-share', 'art-mf-share')`).Execute() + 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")), @@ -309,6 +323,23 @@ var _ = Describe("ShareRepository", func() { Expect(share.Albums).ToNot(ContainElement(HaveField("ID", "art-album-other")), "an album outside the owner's libraries must not appear in the share") }) + + 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(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"))) + Expect(share.Tracks).To(ContainElement(HaveField("ID", "art-ok"))) + Expect(share.Tracks).ToNot(ContainElement(HaveField("ID", "art-other"))) + }) + + It("excludes tracks outside the owner's libraries from a media file share", func() { + 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"))) + }) }) Describe("Ownership Checks", func() { @@ -335,92 +366,114 @@ 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(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(GetDBXBuilder()) + + id, err := attackerRepo.Save(attackerCtx, &model.Share{ + ID: "spoof-save-share", UserID: otherUser.ID, + ResourceType: "media_file", ResourceIDs: "1001", + }) + Expect(err).ToNot(HaveOccurred()) + + 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)) + }) + }) + Describe("Update", 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")) @@ -429,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) @@ -455,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 { @@ -474,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 @@ -485,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)) }) @@ -528,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 { @@ -540,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 9d2ac9590..871f53dd2 100644 --- a/persistence/smart_playlist_repository.go +++ b/persistence/smart_playlist_repository.go @@ -1,11 +1,15 @@ package persistence import ( + "context" + "slices" "time" . "github.com/Masterminds/squirrel" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/utils/slice" + "golang.org/x/text/unicode/norm" ) // PlaylistRepository methods to handle smart playlists, which are defined by criteria and automatically populated @@ -15,65 +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 { - usr := loggedUser(r.ctx) - if !r.shouldRefreshSmartPlaylist(pls, usr) { +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(ctx context.Context, pls *model.Playlist, visited map[string]struct{}) bool { + if _, seen := visited[pls.ID]; seen { + log.Trace(ctx, "Skipping already visited smart playlist", "playlist", pls.Name, "id", pls.ID) + return false + } + visited[pls.ID] = struct{}{} + + 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.Rules, withSmartPlaylistOwner(*usr)) + rulesSQL := newSmartPlaylistCriteria(*pls.NormalizedRules(), withSmartPlaylistOwner(*usr)) - if !r.refreshChildPlaylists(pls, rulesSQL) { + 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 } @@ -81,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 @@ -89,67 +104,86 @@ 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) bool { +func (r *playlistRepository) refreshChildPlaylists(ctx context.Context, pls *model.Playlist, rulesSQL smartPlaylistCriteria, visited map[string]struct{}) bool { childPlaylistIds := rulesSQL.ChildPlaylistIds() - if len(childPlaylistIds) == 0 { + childPlaylistPaths := rulesSQL.ChildPlaylistPaths() + if len(childPlaylistIds) == 0 && len(childPlaylistPaths) == 0 { return true } - childPlaylists, err := r.GetAll(model.QueryOptions{Filters: Eq{"playlist.id": childPlaylistIds}}) + var conditions Or + if len(childPlaylistIds) > 0 { + conditions = append(conditions, Eq{"playlist.id": childPlaylistIds}) + } + if len(childPlaylistPaths) > 0 { + lookupPaths := slices.Concat(slice.Map(childPlaylistPaths, pathVariants)...) + conditions = append(conditions, Eq{"playlist.path": lookupPaths}) + } + + 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, 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 } - found := make(map[string]struct{}, len(childPlaylists)) + found := make(map[string]struct{}, len(childPlaylists)*2) for i := range childPlaylists { found[childPlaylists[i].ID] = struct{}{} - r.refreshSmartPlaylist(&childPlaylists[i]) + if childPlaylists[i].Path != "" { + found[norm.NFC.String(childPlaylists[i].Path)] = struct{}{} + } + 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(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 6f8684d5c..4d4c5edb0 100644 --- a/persistence/smart_playlist_repository_test.go +++ b/persistence/smart_playlist_repository_test.go @@ -1,6 +1,8 @@ package persistence import ( + "context" + "path/filepath" "time" "github.com/navidrome/navidrome/conf" @@ -17,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() { @@ -36,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)) }) @@ -48,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)) @@ -71,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"))) }) }) @@ -85,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)) @@ -100,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)) @@ -124,41 +126,122 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { criteria.Contains{"title": "Day"}, }, } - nestedPls := model.Playlist{Name: "Nested", OwnerID: "userid", Public: true, Rules: childRules} - Expect(repo.Put(&nestedPls)).To(Succeed()) - DeferCleanup(func() { _ = repo.Delete(nestedPls.ID) }) + nestedPls := model.Playlist{Name: "Nested [ID]", OwnerID: "userid", Public: true, Rules: childRules} + Expect(repo.Put(ctx, &nestedPls)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, nestedPls.ID) }) + + childRules = &criteria.Criteria{ + Expression: criteria.All{ + criteria.Eq{"artist": "シートベルツ"}, + }, + } + nestedPathPls := model.Playlist{Name: "Nested [Path]", OwnerID: "userid", Path: "test.nsp", Public: true, Rules: childRules} + 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.All{ + Expression: criteria.Any{ criteria.InPlaylist{"id": nestedPls.ID}, + 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)) // Parent should have tracks from the nested playlist - Expect(pls.Tracks).To(HaveLen(1)) + Expect(pls.Tracks).To(HaveLen(2)) Expect(pls.Tracks[0].MediaFileID).To(Equal(songDayInALife.ID)) - // Nested playlist should now have been refreshed (EvaluatedAt set) - nestedPlsAfterParentGet, err := repo.Get(nestedPls.ID) + // Nested playlists should now have been refreshed (EvaluatedAt set) + 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(ctx, nestedPathPls.ID) Expect(err).ToNot(HaveOccurred()) Expect(nestedPlsAfterParentGet.EvaluatedAt).ToNot(BeNil()) Expect(*nestedPlsAfterParentGet.EvaluatedAt).To(BeTemporally("~", time.Now(), 2*time.Second)) }) }) + It("does not recurse forever when two smart playlists reference each other", func() { + conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second + + plsA := model.Playlist{Name: "Cycle A", OwnerID: "userid", Public: true, Rules: &criteria.Criteria{ + Expression: criteria.All{criteria.Contains{"title": "Day"}}, + }} + 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(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(ctx, &plsA)).To(Succeed()) + + _, err := repo.GetWithTracks(ctx, plsA.ID, true, false) + Expect(err).ToNot(HaveOccurred()) + }) + + It("does not treat an empty path as a reference to every playlist without a path", func() { + conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second + + bystander := model.Playlist{Name: "Bystander", OwnerID: "userid", Public: true, Rules: &criteria.Criteria{ + Expression: criteria.All{criteria.Contains{"title": "Day"}}, + }} + 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(ctx, &parent)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, parent.ID) }) + + _, err := repo.GetWithTracks(ctx, parent.ID, true, false) + Expect(err).ToNot(HaveOccurred()) + + reloaded, err := repo.Get(ctx, bystander.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(reloaded.EvaluatedAt).To(BeNil()) + }) + + It("matches a child path stored in a different Unicode normalization form", func() { + conf.Server.SmartPlaylistRefreshDelay = -1 * time.Second + + 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(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(ctx, &parent)).To(Succeed()) + DeferCleanup(func() { _ = repo.Delete(ctx, parent.ID) }) + + 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)) + }) + When("refresh delay has not expired", func() { It("should NOT refresh tracks for smart playlist referenced in parent smart playlist criteria", func() { conf.Server.SmartPlaylistRefreshDelay = 1 * time.Hour @@ -170,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{ @@ -179,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) @@ -194,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)) @@ -215,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)) @@ -234,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)) @@ -252,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 = "" } }) @@ -262,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") @@ -285,7 +368,7 @@ var _ = Describe("PlaylistRepository - Smart Playlists", func() { AfterEach(func() { if testPlaylistID != "" { - _ = repo.Delete(testPlaylistID) + _ = repo.Delete(ctx, testPlaylistID) testPlaylistID = "" } }) @@ -299,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)) @@ -322,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)) @@ -347,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)) @@ -369,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()) @@ -389,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)) @@ -413,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{ @@ -424,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)) @@ -453,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 }) } @@ -487,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"}) @@ -508,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{ @@ -524,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 @@ -547,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") @@ -568,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") @@ -623,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{ @@ -639,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{ @@ -655,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 @@ -679,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{ @@ -688,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 46ad6a0de..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,26 +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 } - upd := Update(annotationTable).Where(And{ - Eq{annotationTable + ".item_type": r.tableName}, - Eq{annotationTable + ".item_id": prevID}, - }).Set("item_id", newID) - _, err := r.executeSQL(upd) - return err + // 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(ctx, upd); err != nil { + return err + } + // The moved rows change newID's rating population, so its cached average no longer matches + 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 5766f687f..6761331d9 100644 --- a/persistence/sql_annotations_test.go +++ b/persistence/sql_annotations_test.go @@ -15,19 +15,67 @@ 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() { + var prev, next model.Album + + 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(ctx, &prev)).To(Succeed()) + Expect(albumRepo.Put(ctx, &next)).To(Succeed()) + }) + + AfterEach(func() { + _, _ = 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(ctx, 4, prev.ID)).To(Succeed()) + + Expect(albumRepo.ReassignAnnotation(ctx, prev.ID, next.ID)).To(Succeed()) + + 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(ctx, 4, prev.ID)).To(Succeed()) + + Expect(albumRepo.ReassignAnnotation(ctx, prev.ID, next.ID)).To(Succeed()) + + 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(ctx, 4, prev.ID)).To(Succeed()) + Expect(albumRepo.SetRating(ctx, 2, next.ID)).To(Succeed()) + + Expect(albumRepo.ReassignAnnotation(ctx, prev.ID, next.ID)).To(Succeed()) + + got, err := albumRepo.Get(ctx, next.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(got.Rating).To(Equal(2)) + }) }) Describe("annotationBoolFilter", func() { @@ -43,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() { @@ -53,7 +103,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()) @@ -69,7 +119,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()) @@ -82,7 +132,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()) @@ -98,7 +148,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()) @@ -111,14 +161,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()) @@ -135,11 +185,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 { @@ -207,11 +257,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()) @@ -220,16 +270,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()) @@ -237,7 +287,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 5530d2568..03cc6a01b 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,17 +74,33 @@ 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 } +// 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.") @@ -121,10 +136,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 +258,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 +268,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 +285,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 +372,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 } @@ -378,8 +394,9 @@ func wrapCursor[D, T any](cursor iter.Seq2[D, error], toModel func(D) *T) iter.S for row, err := range cursor { m := toModel(row) if m == nil { + // Don't format row: its String() derefs the nil model (golang/go#81238). var zero T - yield(zero, fmt.Errorf("unexpected nil %T (%v): %w", zero, row, err)) + yield(zero, fmt.Errorf("unexpected nil %T: %w", zero, err)) return } if !yield(*m, err) || err != nil { @@ -391,7 +408,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]) } @@ -400,8 +417,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 } @@ -421,7 +438,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]) } @@ -430,28 +447,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 } @@ -468,10 +485,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 } @@ -486,22 +503,19 @@ 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) - if err != nil { - return err - } - if count == 0 { - return r.classifyOwnedWriteMiss(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 @@ -509,13 +523,27 @@ 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 { + 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(id) + return r.classifyOwnedWriteMiss(ctx, rowID) + } + 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 } @@ -523,8 +551,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 } @@ -534,7 +562,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(). @@ -542,22 +570,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 @@ -587,7 +615,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) @@ -595,7 +623,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 } @@ -609,22 +637,30 @@ 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 { - del := Delete(r.tableName).Where(cond) - _, err := r.executeSQL(del) - if errors.Is(err, sql.ErrNoRows) { - return model.ErrNotFound - } +func (r sqlRepository) delete(ctx context.Context, cond Sqlizer) error { + _, err := r.executeSQL(ctx, Delete(r.tableName).Where(cond)) return err } -func (r sqlRepository) logSQL(sql string, args dbx.Params, err error, rowsAffected int64, start time.Time) { +// deleteByID is for single-item deletes that must report a missing row; delete succeeds silently. +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 + } + if count == 0 { + return model.ErrNotFound + } + return nil +} + +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 @@ -634,5 +670,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(ctx) { + log.Warn(append(fields, err)...) + return + } log.Error(append(fields, err)...) } diff --git a/persistence/sql_base_repository_test.go b/persistence/sql_base_repository_test.go index 0f76eb6ab..4b42e7f1d 100644 --- a/persistence/sql_base_repository_test.go +++ b/persistence/sql_base_repository_test.go @@ -13,11 +13,32 @@ 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" }) + 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() { @@ -88,19 +109,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 +133,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 +269,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 +319,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 +333,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 +347,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 +357,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 +380,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 +402,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 19f16b231..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,17 +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.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 } @@ -117,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{ @@ -144,14 +146,21 @@ func (r sqlRepository) GetBookmarks() (model.Bookmarks, error) { return resp, nil } -func (r sqlRepository) cleanBookmarks() 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(ctx, upd) + return err +} + +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 712a928db..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,21 +54,69 @@ 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, userCtx context.Context + var userMr model.MediaFileRepository + + BeforeEach(func() { + adminCtx, otherLib, restrictedUser = restrictedFixture("bmk") + + 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(adminCtx, "bmk-otherlib-track") }) + + userCtx = request.WithUser(ctx, restrictedUser) + userMr = NewMediaFileRepository(GetDBXBuilder()) + }) + + It("does not return bookmarks for tracks outside the user's libraries", func() { + Expect(userMr.AddBookmark(userCtx, "bmk-otherlib-track", "sneaky", 1)).To(Succeed()) + + Expect(userMr.GetBookmarks(userCtx)).To(BeEmpty()) + }) + + It("still returns the bookmark for an admin", func() { + 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(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(userCtx, songAntenna.ID, "allowed", 5)).To(Succeed()) + DeferCleanup(func() { _ = userMr.DeleteBookmark(userCtx, songAntenna.ID) }) + + 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 96fd3efdb..9db56e143 100644 --- a/persistence/transcoding_repository.go +++ b/persistence/transcoding_repository.go @@ -2,7 +2,6 @@ package persistence import ( "context" - "errors" . "github.com/Masterminds/squirrel" "github.com/deluan/rest" @@ -14,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 = "" } @@ -78,50 +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) - id, err := r.put(t.ID, t) - if errors.Is(err, model.ErrNotFound) { - return "", rest.ErrNotFound - } - return id, err + 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) - if errors.Is(err, model.ErrNotFound) { - return rest.ErrNotFound - } + _, 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 } - err := r.delete(Eq{"id": id}) - if errors.Is(err, model.ErrNotFound) { - return rest.ErrNotFound + for _, id := range ids { + if err := r.deleteByID(ctx, id); err != nil { + return err + } } - 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 73250163c..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,78 +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 := 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()) @@ -90,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()) @@ -102,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")) @@ -119,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 9de37876b..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,121 +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 } - usr, err := r.Get(id) - if errors.Is(err, model.ErrNotFound) { - return nil, rest.ErrNotFound - } - return usr, err + 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 } @@ -302,23 +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 } - err := r.Put(u) - if errors.Is(err, model.ErrNotFound) { - return rest.ErrNotFound - } - return err + return r.Put(ctx, u) } func validatePasswordChange(newUser *model.User, logged *model.User) error { @@ -347,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 } @@ -388,22 +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 } - err := r.delete(Eq{"id": id}) - if errors.Is(err, model.ErrNotFound) { - return rest.ErrNotFound - } - if 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 } @@ -413,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 @@ -422,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 { @@ -437,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 { @@ -463,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 @@ -472,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 @@ -483,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 } } @@ -504,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"). @@ -512,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 } @@ -529,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 dc519d0a1..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("")) }) @@ -230,31 +232,37 @@ 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(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 { @@ -265,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")), @@ -282,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) @@ -303,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")) }) }) @@ -322,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)) @@ -366,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)) @@ -391,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)) }) @@ -414,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) @@ -425,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() { @@ -447,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 @@ -473,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 @@ -508,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)) @@ -529,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{ @@ -546,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)) @@ -583,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 @@ -599,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)) @@ -617,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) @@ -678,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}")) @@ -697,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)) }) @@ -727,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")) @@ -746,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) @@ -770,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/README.md b/plugins/README.md index 7042b8c45..7dca3a5f2 100644 --- a/plugins/README.md +++ b/plugins/README.md @@ -401,7 +401,7 @@ import "github.com/navidrome/navidrome/plugins/pdk/go/host" ### HTTP -Make HTTP requests to external services. This is a dedicated host service (separate from Extism's built-in HTTP support) with additional features like timeouts and redirect control. +Make HTTP requests to external services, with timeouts, redirect control, and protection against reaching private network addresses. This is the only supported way to make HTTP requests: Extism's built-in HTTP is disabled (`pdk.NewHTTPRequest` in Go, `http::request` in Rust, `extism.Http.request` in Python, `Http.request` in JS). **Manifest permission:** @@ -416,6 +416,8 @@ Make HTTP requests to external services. This is a dedicated host service (separ } ``` +**Private addresses:** the check runs on the resolved IP when connecting. A named host entry (`api.example.com`, `*.example.com`) never authorizes a loopback, private or link-local address on its own, even if its DNS points there. To reach a service on the local network, also list its IP or a CIDR (`192.168.1.10`, `10.0.0.0/8`), or use `"*"` when the user configures the address. Without `requiredHosts`, only public addresses are allowed. + **Host functions:** | Function | Parameters | Returns | @@ -700,6 +702,8 @@ Establish persistent WebSocket connections to external services. Your plugin mus } ``` +`requiredHosts` is mandatory here: leave it out and every connection is blocked. Unlike HTTP, there is no fallback to public addresses. Entries follow the same [private address rules](#http). + **Host functions:** | Function | Parameters | Description | @@ -1242,6 +1246,8 @@ extism-py plugin.wasm -o plugin.wasm *.py zip -j my-plugin.ndp manifest.json plugin.wasm ``` +There is no Python PDK, so call host services directly: import them from the `extism:host/user` namespace (e.g. `http_send`) with `@extism.import_fn`, and exchange JSON through Extism memory. Each function takes a JSON request and returns a JSON response with an `error` field on failure. For HTTP, send `{"request": {"method": "GET", "url": "..."}}` and read `result.statusCode` and `result.body` (base64). See [coverartarchive-py](examples/coverartarchive-py/) and [nowplaying-py](examples/nowplaying-py/). + ### Using XTP CLI (Scaffolding) Bootstrap a new plugin from a schema: diff --git a/plugins/cmd/ndpgen/internal/generator_test.go b/plugins/cmd/ndpgen/internal/generator_test.go index ea94c7fc6..555906489 100644 --- a/plugins/cmd/ndpgen/internal/generator_test.go +++ b/plugins/cmd/ndpgen/internal/generator_test.go @@ -1790,6 +1790,33 @@ var _ = Describe("Rust Generation", func() { Expect(codeStr).NotTo(ContainSubstring("return args.Get(0).(*HTTPRequest)")) }) }) + + Describe("Deprecated PDK functions", func() { + symbols := &PDKSymbols{ + Functions: []PDKFunc{ + { + Name: "NewHTTPRequest", + Doc: "NewHTTPRequest returns a new `HTTPRequest`.", + Params: []PDKParam{{Name: "method", Type: "HTTPMethod"}, {Name: "url", Type: "string"}}, + Returns: []PDKReturn{{Type: "*HTTPRequest"}}, + Deprecated: "Use host.HTTPSend instead.", + }, + }, + } + + DescribeTable("emits a Deprecated paragraph after the doc line", + func(generate func(*PDKSymbols) ([]byte, error)) { + code, err := generate(symbols) + Expect(err).NotTo(HaveOccurred()) + _, err = format.Source(code) + Expect(err).NotTo(HaveOccurred()) + Expect(string(code)).To(ContainSubstring( + "// NewHTTPRequest NewHTTPRequest returns a new `HTTPRequest`.\n//\n// Deprecated: Use host.HTTPSend instead.\nfunc NewHTTPRequest(")) + }, + Entry("WASM wrapper", GeneratePDKGo), + Entry("native stub", GeneratePDKGoStub), + ) + }) }) }) diff --git a/plugins/cmd/ndpgen/internal/pdk_parser.go b/plugins/cmd/ndpgen/internal/pdk_parser.go index 4756334dd..7607fe881 100644 --- a/plugins/cmd/ndpgen/internal/pdk_parser.go +++ b/plugins/cmd/ndpgen/internal/pdk_parser.go @@ -50,6 +50,12 @@ type PDKFunc struct { Params []PDKParam Returns []PDKReturn IsVariadic bool + Deprecated string // Rendered as a "Deprecated:" paragraph when set +} + +// deprecatedPDKFuncs marks extism functions that do not work inside Navidrome, with what to use instead. +var deprecatedPDKFuncs = map[string]string{ + "NewHTTPRequest": "Navidrome does not enable extism's http_request host function, so every request sent this way fails. Use host.HTTPSend instead.", } // PDKParam represents a function parameter. @@ -156,6 +162,7 @@ func ParseExtismPDK() (*PDKSymbols, error) { t.Methods = append(t.Methods, fn) } } else { + fn.Deprecated = deprecatedPDKFuncs[fn.Name] symbols.Functions = append(symbols.Functions, fn) } } diff --git a/plugins/cmd/ndpgen/internal/templates/client.rs.tmpl b/plugins/cmd/ndpgen/internal/templates/client.rs.tmpl index f8b786849..786a50db5 100644 --- a/plugins/cmd/ndpgen/internal/templates/client.rs.tmpl +++ b/plugins/cmd/ndpgen/internal/templates/client.rs.tmpl @@ -11,7 +11,7 @@ use serde::{Deserialize, Serialize}; {{if .Doc}} {{rustDocComment .Doc}} {{else}} -{{end}}#[derive(Debug, Clone, Serialize, Deserialize)] +{{end}}#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct {{.Name}} { {{- range .Fields}} diff --git a/plugins/cmd/ndpgen/internal/templates/pdk.go.tmpl b/plugins/cmd/ndpgen/internal/templates/pdk.go.tmpl index ebaf88df0..117bd9232 100644 --- a/plugins/cmd/ndpgen/internal/templates/pdk.go.tmpl +++ b/plugins/cmd/ndpgen/internal/templates/pdk.go.tmpl @@ -40,6 +40,10 @@ const ( {{- if .Doc}} // {{.Name}} {{firstSentence .Doc}} {{- end}} +{{- if .Deprecated}} +// +// Deprecated: {{.Deprecated}} +{{- end}} func {{.Name}}({{paramList .Params}}){{returnList .Returns}} { {{- if .Returns}} return extism.{{.Name}}({{argList .Params}}) diff --git a/plugins/cmd/ndpgen/internal/templates/pdk_stub.go.tmpl b/plugins/cmd/ndpgen/internal/templates/pdk_stub.go.tmpl index 2bd73f107..27d06faa0 100644 --- a/plugins/cmd/ndpgen/internal/templates/pdk_stub.go.tmpl +++ b/plugins/cmd/ndpgen/internal/templates/pdk_stub.go.tmpl @@ -31,6 +31,10 @@ func ResetMock() { {{- if .Doc}} // {{.Name}} {{firstSentence .Doc}} {{- end}} +{{- if .Deprecated}} +// +// Deprecated: {{.Deprecated}} +{{- end}} func {{.Name}}({{paramList .Params}}){{returnList .Returns}} { {{- if .Returns}} args := PDKMock.Called({{argList .Params}}) diff --git a/plugins/cmd/ndpgen/testdata/comprehensive_client_expected.rs b/plugins/cmd/ndpgen/testdata/comprehensive_client_expected.rs index 08dae2901..1bd953640 100644 --- a/plugins/cmd/ndpgen/testdata/comprehensive_client_expected.rs +++ b/plugins/cmd/ndpgen/testdata/comprehensive_client_expected.rs @@ -29,14 +29,14 @@ mod base64_bytes { } } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct User2 { pub id: String, pub name: String, } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct Filter2 { pub active: bool, diff --git a/plugins/cmd/ndpgen/testdata/list_client_expected.rs b/plugins/cmd/ndpgen/testdata/list_client_expected.rs index 9b54f7544..5227cadf4 100644 --- a/plugins/cmd/ndpgen/testdata/list_client_expected.rs +++ b/plugins/cmd/ndpgen/testdata/list_client_expected.rs @@ -6,7 +6,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct Filter { pub active: bool, diff --git a/plugins/cmd/ndpgen/testdata/search_client_expected.rs b/plugins/cmd/ndpgen/testdata/search_client_expected.rs index b0ab2505a..396a22b4e 100644 --- a/plugins/cmd/ndpgen/testdata/search_client_expected.rs +++ b/plugins/cmd/ndpgen/testdata/search_client_expected.rs @@ -6,7 +6,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct Result { pub id: String, diff --git a/plugins/cmd/ndpgen/testdata/store_client_expected.rs b/plugins/cmd/ndpgen/testdata/store_client_expected.rs index 25d2af2e0..3b9d5ac17 100644 --- a/plugins/cmd/ndpgen/testdata/store_client_expected.rs +++ b/plugins/cmd/ndpgen/testdata/store_client_expected.rs @@ -6,7 +6,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct Item { pub id: String, diff --git a/plugins/cmd/ndpgen/testdata/users_client_expected.rs b/plugins/cmd/ndpgen/testdata/users_client_expected.rs index 40daa9cfd..4bac43ce8 100644 --- a/plugins/cmd/ndpgen/testdata/users_client_expected.rs +++ b/plugins/cmd/ndpgen/testdata/users_client_expected.rs @@ -6,7 +6,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct User { pub id: String, diff --git a/plugins/examples/coverartarchive-py/Makefile b/plugins/examples/coverartarchive-py/Makefile index e3cd60d1c..c47c242d7 100644 --- a/plugins/examples/coverartarchive-py/Makefile +++ b/plugins/examples/coverartarchive-py/Makefile @@ -1,5 +1,5 @@ # Build the Cover Art Archive Python plugin -.PHONY: build test clean +.PHONY: build clean WASM_FILE = coverartarchive-py.wasm @@ -8,20 +8,5 @@ build: $(WASM_FILE) $(WASM_FILE): plugin/__init__.py extism-py plugin/__init__.py -o $(WASM_FILE) -test: build - @echo "Testing nd_manifest..." - extism call $(WASM_FILE) nd_manifest --wasi - @echo "" - @echo "Testing nd_get_album_images with Portishead's Dummy MBID..." - extism call $(WASM_FILE) nd_get_album_images --wasi \ - --input '{"name":"Dummy","artist":"Portishead","mbid":"76df3287-6cda-33eb-8e9a-044b5e15ffdd"}' \ - --allow-host "coverartarchive.org" --allow-host "archive.org" - -test-error: build - @echo "Testing error case (missing MBID)..." - -extism call $(WASM_FILE) nd_get_album_images --wasi \ - --input '{"name":"Test Album","artist":"Test Artist"}' \ - --allow-host "coverartarchive.org" - clean: rm -f $(WASM_FILE) diff --git a/plugins/examples/coverartarchive-py/README.md b/plugins/examples/coverartarchive-py/README.md index 77957ac4f..c779f4a43 100644 --- a/plugins/examples/coverartarchive-py/README.md +++ b/plugins/examples/coverartarchive-py/README.md @@ -51,14 +51,7 @@ zip -j coverartarchive-py.ndp manifest.json plugin.wasm ## Testing -Extract the wasm file and test: - -```bash -unzip -p coverartarchive-py.ndp plugin.wasm > coverartarchive-py.wasm -extism call coverartarchive-py.wasm nd_get_album_images --wasi \ - --input '{"name":"Dummy","artist":"Portishead","mbid":"76df3287-6cda-33eb-8e9a-044b5e15ffdd"}' \ - --allow-host "coverartarchive.org" --allow-host "archive.org" -``` +The plugin makes HTTP requests through Navidrome's `http_send` host function, so it only runs inside Navidrome (the `extism` CLI can't provide that function). Install the `.ndp` as described above, then open an album that has a MusicBrainz ID and check the Navidrome logs. ## How It Works diff --git a/plugins/examples/coverartarchive-py/plugin/__init__.py b/plugins/examples/coverartarchive-py/plugin/__init__.py index 3c1d4149e..c4a0480b5 100644 --- a/plugins/examples/coverartarchive-py/plugin/__init__.py +++ b/plugins/examples/coverartarchive-py/plugin/__init__.py @@ -5,16 +5,27 @@ # # Build with: # extism-py plugin/__init__.py -o coverartarchive-py.wasm -# -# Test with: -# extism call coverartarchive-py.wasm nd_get_album_images --wasi \ -# --input '{"name":"Dummy","artist":"Portishead","mbid":"76df3287-6cda-33eb-8e9a-044b5e15ffdd"}' \ -# --allow-host "coverartarchive.org" --allow-host "archive.org" +import base64 import extism import json +@extism.import_fn("extism:host/user", "http_send") +def _http_send(offset: int) -> int: ... + + +def http_get(url): + """GET url via Navidrome's HTTP host service. Returns (status_code, body_bytes).""" + request = json.dumps({"request": {"method": "GET", "url": url}}).encode("utf-8") + response_offset = _http_send(extism.memory.alloc(request).offset) + resp = json.loads(extism.memory.string(extism.memory.find(response_offset))) + if resp.get("error"): + raise Exception(f"HTTP request failed: {resp['error']}") + result = resp["result"] + return result["statusCode"], base64.b64decode(result.get("body", "")) + + @extism.plugin_fn def nd_get_album_images(): """Retrieve album cover images from Cover Art Archive.""" @@ -26,13 +37,13 @@ def nd_get_album_images(): # Query Cover Art Archive API url = f"https://coverartarchive.org/release/{mbid}" - response = extism.Http.request(url, meth="GET") - - if response.status_code != 200: - raise Exception(f"not found: CAA returned status {response.status_code}") - + status, body = http_get(url) + + if status != 200: + raise Exception(f"not found: CAA returned status {status}") + try: - data = json.loads(response.data_str()) + data = json.loads(body) except json.JSONDecodeError: raise Exception("not found: invalid JSON response") diff --git a/plugins/examples/discord-rich-presence-rs/src/rpc.rs b/plugins/examples/discord-rich-presence-rs/src/rpc.rs index 3de9eff63..fa72f7466 100644 --- a/plugins/examples/discord-rich-presence-rs/src/rpc.rs +++ b/plugins/examples/discord-rich-presence-rs/src/rpc.rs @@ -4,7 +4,7 @@ //! presence updates, and heartbeat management. use extism_pdk::*; -use nd_pdk::host::{cache, scheduler, websocket}; +use nd_pdk::host::{cache, http, scheduler, websocket}; use serde::{Deserialize, Serialize}; // ============================================================================ @@ -359,19 +359,32 @@ fn find_username_for_connection(connection_id: &str) -> Result, E Ok(cache::get_string(&reverse_key)?.filter(|s| !s.is_empty())) } -fn get_discord_gateway() -> Result { - let req = HttpRequest::new("https://discord.com/api/gateway") - .with_method("GET"); +fn send_http( + method: &str, + url: &str, + headers: std::collections::HashMap, + body: Vec, +) -> Result { + http::send(http::HTTPRequest { + method: method.into(), + url: url.into(), + headers, + body, + ..Default::default() + })? + .ok_or_else(|| Error::msg("empty HTTP response")) +} - let resp = http::request::(&req, None::)?; - if resp.status_code() >= 400 { +fn get_discord_gateway() -> Result { + let resp = send_http("GET", "https://discord.com/api/gateway", Default::default(), Vec::new())?; + if resp.status_code >= 400 { return Err(Error::msg(format!( "Failed to get Discord gateway: HTTP {}", - resp.status_code() + resp.status_code ))); } - let body = resp.body(); + let body = resp.body; let data: std::collections::HashMap = serde_json::from_slice(&body) .map_err(|e| Error::msg(format!("Failed to parse gateway response: {}", e)))?; @@ -487,23 +500,22 @@ fn process_image_inner( client_id ); - let req = HttpRequest::new(&api_url) - .with_method("POST") - .with_header("Authorization", token) - .with_header("Content-Type", "application/json"); - - let resp = http::request::(&req, Some(body))?; - if resp.status_code() >= 400 { + let headers = std::collections::HashMap::from([ + ("Authorization".to_string(), token.to_string()), + ("Content-Type".to_string(), "application/json".to_string()), + ]); + let resp = send_http("POST", &api_url, headers, body.into_bytes())?; + if resp.status_code >= 400 { if is_default { return Err(Error::msg(format!( "failed to process default image: HTTP {}", - resp.status_code() + resp.status_code ))); } return process_image_inner(DEFAULT_IMAGE, client_id, token, true); } - let body = resp.body(); + let body = resp.body; let data: Vec> = serde_json::from_slice(&body) .map_err(|e| Error::msg(format!("Failed to parse image response: {}", e)))?; diff --git a/plugins/examples/webhook-rs/src/lib.rs b/plugins/examples/webhook-rs/src/lib.rs index 743c03744..5f5d972a0 100644 --- a/plugins/examples/webhook-rs/src/lib.rs +++ b/plugins/examples/webhook-rs/src/lib.rs @@ -12,7 +12,8 @@ //! urls = "https://example.com/webhook1,https://example.com/webhook2" //! ``` -use extism_pdk::{config, error, http, info, warn, HttpRequest}; +use extism_pdk::{config, error, info, warn}; +use nd_pdk::host::http::{self, HTTPRequest}; use nd_pdk::scrobbler::{ Error, IsAuthorizedRequest, NowPlayingRequest, PlaybackReportRequest, ScrobbleRequest, Scrobbler, @@ -90,11 +91,15 @@ impl Scrobbler for WebhookPlugin { let full_url = format!("{}{}", url, query); info!("Sending webhook to: {}", full_url); - let http_req = HttpRequest::new(&full_url); - match http::request::<()>(&http_req, None) { + let http_req = HTTPRequest { + method: "GET".into(), + url: full_url, + ..Default::default() + }; + match http::send(http_req) { Ok(res) => { - let status = res.status_code(); - if status >= 200 && status < 300 { + let status = res.map_or(0, |r| r.status_code); + if (200..300).contains(&status) { info!("Webhook succeeded: {} (status {})", url, status); } else { warn!("Webhook returned non-2xx status: {} (status {})", url, status); diff --git a/plugins/host_httpclient.go b/plugins/host_httpclient.go index d52898bdd..6a9a4d2d6 100644 --- a/plugins/host_httpclient.go +++ b/plugins/host_httpclient.go @@ -10,11 +10,13 @@ import ( "net/http" "net/url" "strings" + "syscall" "time" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/plugins/host" "github.com/navidrome/navidrome/utils/httpclient" + "github.com/navidrome/navidrome/utils/netguard" ) const ( @@ -34,6 +36,7 @@ type httpServiceImpl struct { pluginName string requiredHosts []string client *http.Client + transport *http.Transport } // newHTTPService creates a new HTTPService for a plugin. @@ -46,8 +49,15 @@ func newHTTPService(pluginName string, permission *HTTPPermission) *httpServiceI pluginName: pluginName, requiredHosts: requiredHosts, } + svc.transport = http.DefaultTransport.(*http.Transport).Clone() + svc.transport.DialContext = (&net.Dialer{ + 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 = httpclient.New(0) + svc.client = &http.Client{Transport: httpclient.NewTransport(svc.transport)} svc.client.CheckRedirect = func(req *http.Request, via []*http.Request) error { if req.Context().Value(noFollowRedirectsKey) != nil { return http.ErrUseLastResponse @@ -65,6 +75,12 @@ func newHTTPService(pluginName string, permission *HTTPPermission) *httpServiceI return svc } +// Close releases the plugin's pooled connections when the plugin is unloaded. +func (s *httpServiceImpl) Close() error { + s.transport.CloseIdleConnections() + return nil +} + func (s *httpServiceImpl) Send(ctx context.Context, request host.HTTPRequest) (*host.HTTPResponse, error) { // Parse and validate URL parsedURL, err := url.Parse(request.URL) @@ -145,7 +161,7 @@ func (s *httpServiceImpl) validateHost(ctx context.Context, hostStr string) erro hostname := extractHostname(hostStr) if len(s.requiredHosts) > 0 { - if !s.isHostAllowed(hostname) { + if !isHostInAllowlist(s.requiredHosts, hostname) { return fmt.Errorf("host %q is not allowed", hostStr) } return nil @@ -159,13 +175,8 @@ func (s *httpServiceImpl) validateHost(ctx context.Context, hostStr string) erro return nil } -func (s *httpServiceImpl) isHostAllowed(hostname string) bool { - for _, pattern := range s.requiredHosts { - if matchHostPattern(pattern, hostname) { - return true - } - } - return false +func (s *httpServiceImpl) dialControl(_, address string, _ syscall.RawConn) error { + return checkPrivateDial(s.requiredHosts, address) } // extractHostname returns the hostname portion of a host string, stripping @@ -182,11 +193,8 @@ func extractHostname(hostStr string) string { return hostStr } -// isPrivateOrLoopback returns true if the given hostname resolves to or is -// a private, loopback, or link-local IP address. This includes: -// IPv4: 127.0.0.0/8, 10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 169.254.0.0/16 -// IPv6: ::1, fc00::/7, fe80::/10 -// It also blocks "localhost" by name. +// isPrivateOrLoopback is a pre-flight check on the literal host (IP or "localhost"); it does not +// resolve names, so dialControl remains the real guard. func isPrivateOrLoopback(hostname string) bool { if strings.EqualFold(hostname, "localhost") { return true @@ -195,7 +203,7 @@ func isPrivateOrLoopback(hostname string) bool { if ip == nil { return false } - return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() + return netguard.IsPrivateIP(ip) } // Verify interface implementation diff --git a/plugins/host_httpclient_test.go b/plugins/host_httpclient_test.go index 4eb34a247..81e3192fd 100644 --- a/plugins/host_httpclient_test.go +++ b/plugins/host_httpclient_test.go @@ -3,6 +3,7 @@ package plugins import ( "context" "io" + "net" "net/http" "net/http/httptest" "strings" @@ -19,6 +20,10 @@ var _ = Describe("httpServiceImpl", func() { ts *httptest.Server ) + BeforeEach(func() { + stubLocalhostDNS() + }) + AfterEach(func() { if ts != nil { ts.Close() @@ -43,6 +48,21 @@ var _ = Describe("httpServiceImpl", func() { Expect(err.Error()).To(ContainSubstring("private/loopback")) }) + It("should block a symbolic hostname that resolves to loopback (SSRF)", func() { + ts = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(200) + })) + // The trailing dot passes the pre-flight string check; only the dial-time guard catches it. + _, port, _ := net.SplitHostPort(strings.TrimPrefix(ts.URL, "http://")) + _, err := svc.Send(context.Background(), host.HTTPRequest{ + Method: "GET", + URL: "http://localhost.:" + port + "/test", + TimeoutMs: 1000, + }) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("private/loopback")) + }) + It("should block requests to localhost by name", func() { _, err := svc.Send(context.Background(), host.HTTPRequest{ Method: "GET", @@ -419,13 +439,53 @@ var _ = Describe("httpServiceImpl", func() { Expect(resp).To(BeNil()) }) + It("blocks a private IP reached via a hostname allowlist entry (rebinding protection)", func() { + // Allowlisting a name authorizes the external service, not whatever private + // IP it may resolve or rebind to. Only literal IP/CIDR entries do that. + svc.requiredHosts = []string{"api.example.com"} + err := svc.dialControl("tcp", "10.0.0.1:80", nil) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("private/loopback")) + }) + + It("allows a private IP that an explicit CIDR allowlist entry authorizes", func() { + svc.requiredHosts = []string{"10.0.0.0/8"} + Expect(svc.dialControl("tcp", "10.0.0.1:80", nil)).To(Succeed()) + }) + + It("allows private IPs when the allowlist is the bare '*' wildcard", func() { + svc.requiredHosts = []string{"*"} + Expect(svc.dialControl("tcp", "192.168.1.10:8000", nil)).To(Succeed()) + Expect(svc.dialControl("tcp", "127.0.0.1:8000", nil)).To(Succeed()) + }) + + It("still blocks private IPs for a subdomain wildcard entry", func() { + svc.requiredHosts = []string{"*.example.com"} + Expect(svc.dialControl("tcp", "10.0.0.1:80", nil)).To(MatchError(ContainSubstring("private/loopback"))) + }) + + It("closes idle pooled connections on Close", func() { + closed := make(chan struct{}) + ts = httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})) + ts.Config.ConnState = func(_ net.Conn, state http.ConnState) { + if state == http.StateClosed { + close(closed) + } + } + ts.Start() + svc.requiredHosts = []string{"127.0.0.1"} + _, err := svc.Send(context.Background(), host.HTTPRequest{Method: "GET", URL: ts.URL, TimeoutMs: 1000}) + Expect(err).ToNot(HaveOccurred()) + Expect(svc.Close()).To(Succeed()) + Eventually(closed).Should(BeClosed()) + }) + It("should allow wildcard host patterns", func() { ts = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte("wildcard")) })) - // *.allowed.org is in the requiredHosts from BeforeEach, but test server is 127.0.0.1 - // Override with a wildcard that matches the test server - svc.requiredHosts = []string{"*.0.0.1"} + // The literal IP is what authorizes the loopback dial under the private-IP guard. + svc.requiredHosts = []string{"*.0.0.1", "127.0.0.1"} resp, err := svc.Send(context.Background(), host.HTTPRequest{ Method: "GET", URL: ts.URL, @@ -566,6 +626,11 @@ var _ = Describe("isPrivateOrLoopback", func() { Expect(isPrivateOrLoopback("fe80::1")).To(BeTrue()) }) + It("should detect unspecified addresses, which dial the local host", func() { + Expect(isPrivateOrLoopback("0.0.0.0")).To(BeTrue()) + Expect(isPrivateOrLoopback("::")).To(BeTrue()) + }) + It("should allow public IPs", func() { Expect(isPrivateOrLoopback("8.8.8.8")).To(BeFalse()) Expect(isPrivateOrLoopback("203.0.113.1")).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_netguard.go b/plugins/host_netguard.go new file mode 100644 index 000000000..ab092009a --- /dev/null +++ b/plugins/host_netguard.go @@ -0,0 +1,58 @@ +package plugins + +import ( + "fmt" + "net" + "slices" + + "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 { + if slices.Contains(requiredHosts, "*") { + return nil + } + host, _, err := net.SplitHostPort(address) + if err != nil { + return err + } + ip := net.ParseIP(host) + if ip == nil || !netguard.IsPrivateIP(ip) { + return nil + } + for _, entry := range requiredHosts { + if ipMatchesEntry(entry, ip) { + return nil + } + } + return fmt.Errorf("dial to private/loopback address %q blocked: requires an explicit IP or CIDR in requiredHosts", address) +} + +func isHostInAllowlist(requiredHosts []string, hostname string) bool { + ip := net.ParseIP(hostname) + for _, pattern := range requiredHosts { + if matchHostPattern(pattern, hostname) { + return true + } + if ip != nil && ipMatchesEntry(pattern, ip) { + return true + } + } + return false +} + +// ipMatchesEntry reports whether a requiredHosts entry is a literal IP or CIDR that covers ip. +func ipMatchesEntry(entry string, ip net.IP) bool { + if _, cidr, err := net.ParseCIDR(entry); err == nil { + return cidr.Contains(ip) + } + if entryIP := net.ParseIP(entry); entryIP != nil { + return entryIP.Equal(ip) + } + return false +} 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_taskqueue_test.go b/plugins/host_taskqueue_test.go index 72fb4cccd..993ddc7f6 100644 --- a/plugins/host_taskqueue_test.go +++ b/plugins/host_taskqueue_test.go @@ -689,11 +689,12 @@ var _ = Describe("TaskQueueService", func() { copy(times, dispatchTimes) mu.Unlock() - // Consecutive dispatches should have at least ~160ms gap (80% of 200ms) + // Wake-up latency varies per worker, so check offsets from the first dispatch, not gaps. for i := 1; i < len(times); i++ { - gap := times[i].Sub(times[i-1]) - Expect(gap).To(BeNumerically(">=", 160*time.Millisecond), - fmt.Sprintf("gap between dispatch %d and %d was %v, expected >= 160ms", i-1, i, gap)) + offset := times[i].Sub(times[0]) + minOffset := time.Duration(i)*200*time.Millisecond - 50*time.Millisecond + Expect(offset).To(BeNumerically(">=", minOffset), + fmt.Sprintf("dispatch %d ran %v after the first, expected >= %v", i, offset, minOffset)) } }) }) 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 82aded0cb..a58deb129 100644 --- a/plugins/host_websocket.go +++ b/plugins/host_websocket.go @@ -5,10 +5,12 @@ import ( "errors" "fmt" "maps" + "net" "net/http" "net/url" "strings" "sync" + "syscall" "time" "github.com/gorilla/websocket" @@ -54,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 @@ -112,6 +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, Resolver: dialResolver}).DialContext, } conn, resp, err := dialer.DialContext(ctx, urlStr, httpHeaders) @@ -243,14 +246,11 @@ func (s *webSocketServiceImpl) getConnection(connectionID string) (*wsConnection } func (s *webSocketServiceImpl) isHostAllowed(host string) bool { - hostWithoutPort := extractHostname(host) + return isHostInAllowlist(s.requiredHosts, extractHostname(host)) +} - for _, pattern := range s.requiredHosts { - if matchHostPattern(pattern, hostWithoutPort) { - return true - } - } - return false +func (s *webSocketServiceImpl) dialControl(_, address string, _ syscall.RawConn) error { + return checkPrivateDial(s.requiredHosts, address) } // matchHostPattern matches a host against a pattern. diff --git a/plugins/host_websocket_test.go b/plugins/host_websocket_test.go index 9f2d20bef..9b94e4b70 100644 --- a/plugins/host_websocket_test.go +++ b/plugins/host_websocket_test.go @@ -8,6 +8,7 @@ import ( "maps" "net/http" "net/http/httptest" + "net/url" "os" "path/filepath" "strings" @@ -499,6 +500,51 @@ var _ = Describe("WebSocketService", Ordered, func() { }) }) + Describe("Private address protection", func() { + var wsServer *httptest.Server + 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) { + if conn, err := upgrader.Upgrade(w, r, nil); err == nil { + _, _, _ = conn.ReadMessage() + } + })) + }) + + AfterEach(func() { + testService.closeAllConnections() + testService.requiredHosts = savedHosts + wsServer.Close() + }) + + serverPort := func() string { + u, _ := url.Parse(wsServer.URL) + return u.Port() + } + + It("blocks an allowlisted hostname that resolves to loopback", func() { + testService.requiredHosts = []string{"localhost."} + _, err := testService.Connect(GinkgoT().Context(), "ws://localhost.:"+serverPort(), nil, "") + Expect(err).To(MatchError(ContainSubstring("private/loopback"))) + }) + + It("allows loopback when a CIDR entry covers it", func() { + testService.requiredHosts = []string{"127.0.0.0/8"} + _, err := testService.Connect(GinkgoT().Context(), "ws://127.0.0.1:"+serverPort(), nil, "") + Expect(err).ToNot(HaveOccurred()) + }) + + It("allows loopback when the allowlist is the bare '*' wildcard", func() { + testService.requiredHosts = []string{"*"} + _, err := testService.Connect(GinkgoT().Context(), "ws://localhost.:"+serverPort(), nil, "") + Expect(err).ToNot(HaveOccurred()) + }) + }) + Describe("Plugin Unload", func() { It("should close all connections when plugin is unloaded", func() { // Create a fresh server for this test 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 46da56396..cca87e5b0 100644 --- a/plugins/manager_loader.go +++ b/plugins/manager_loader.go @@ -148,7 +148,7 @@ var hostServices = []hostServiceEntry{ create: func(ctx *serviceContext) ([]extism.HostFunction, io.Closer, error) { perm := ctx.permissions.Http service := newHTTPService(ctx.pluginName, perm) - return host.RegisterHTTPHostFunctions(service), nil, nil + return host.RegisterHTTPHostFunctions(service), service, nil }, }, { @@ -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) } @@ -498,18 +498,12 @@ func parsePluginConfig(configJSON string) (map[string]string, error) { return pluginConfig, nil } -// buildExtismManifest describes the plugin to extism. It must never set -// AllowedPaths: extism would replace our jailed FSConfig with plain dir mounts. +// buildExtismManifest describes the plugin to extism. It must never set AllowedPaths (extism would replace our +// jailed FSConfig) nor AllowedHosts (extism's http_request has no SSRF guard; plugins must use host.HTTPSend). func buildExtismManifest(pkg *ndpPackage, pluginConfig map[string]string) extism.Manifest { - manifest := extism.Manifest{ + return extism.Manifest{ Wasm: []extism.Wasm{extism.WasmData{Data: pkg.WasmBytes, Name: "main"}}, Config: pluginConfig, Timeout: uint64(defaultTimeout.Milliseconds()), } - if pkg.Manifest.Permissions != nil && pkg.Manifest.Permissions.Http != nil { - if hosts := pkg.Manifest.Permissions.Http.RequiredHosts; len(hosts) > 0 { - manifest.AllowedHosts = hosts - } - } - return manifest } diff --git a/plugins/manager_loader_test.go b/plugins/manager_loader_test.go index 6326b2218..85a636202 100644 --- a/plugins/manager_loader_test.go +++ b/plugins/manager_loader_test.go @@ -23,8 +23,8 @@ var _ = Describe("buildExtismManifest", func() { Expect(buildExtismManifest(pkg, nil).AllowedPaths).To(BeEmpty()) }) - It("carries the hosts the plugin is allowed to reach", func() { - Expect(buildExtismManifest(pkg, nil).AllowedHosts).To(Equal([]string{"example.com"})) + It("never sets AllowedHosts, so plugin HTTP can't bypass the host service's SSRF guard", func() { + Expect(buildExtismManifest(pkg, nil).AllowedHosts).To(BeEmpty()) }) }) 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/plugins/manifest-schema.json b/plugins/manifest-schema.json index dae0fd937..056a1caed 100644 --- a/plugins/manifest-schema.json +++ b/plugins/manifest-schema.json @@ -134,7 +134,7 @@ }, "requiredHosts": { "type": "array", - "description": "List of required host patterns for HTTP requests (e.g., 'api.example.com', '*.musicbrainz.org')", + "description": "List of required host patterns for HTTP requests (e.g., 'api.example.com', '*.musicbrainz.org'). A named host alone can't reach a private address; also list an IP, a CIDR (e.g., '10.0.0.0/8') or '*' for that", "items": { "type": "string" } @@ -174,7 +174,7 @@ }, "requiredHosts": { "type": "array", - "description": "List of required host patterns for WebSocket connections (e.g., 'api.example.com', '*.musicbrainz.org')", + "description": "List of required host patterns for WebSocket connections (e.g., 'api.example.com', '*.musicbrainz.org'). Required: with no entries every connection is blocked. A named host alone can't reach a private address; also list an IP, a CIDR (e.g., '10.0.0.0/8') or '*' for that", "items": { "type": "string" } diff --git a/plugins/manifest_gen.go b/plugins/manifest_gen.go index 6f5c01596..b5e9afd96 100644 --- a/plugins/manifest_gen.go +++ b/plugins/manifest_gen.go @@ -51,7 +51,8 @@ type HTTPPermission struct { Reason *string `json:"reason,omitempty" yaml:"reason,omitempty" mapstructure:"reason,omitempty"` // List of required host patterns for HTTP requests (e.g., 'api.example.com', - // '*.musicbrainz.org') + // '*.musicbrainz.org'). A named host alone can't reach a private address; also + // list an IP, a CIDR (e.g., '10.0.0.0/8') or '*' for that RequiredHosts []string `json:"requiredHosts,omitempty" yaml:"requiredHosts,omitempty" mapstructure:"requiredHosts,omitempty"` } @@ -264,6 +265,8 @@ type WebSocketPermission struct { Reason *string `json:"reason,omitempty" yaml:"reason,omitempty" mapstructure:"reason,omitempty"` // List of required host patterns for WebSocket connections (e.g., - // 'api.example.com', '*.musicbrainz.org') + // 'api.example.com', '*.musicbrainz.org'). Required: with no entries every + // connection is blocked. A named host alone can't reach a private address; also + // list an IP, a CIDR (e.g., '10.0.0.0/8') or '*' for that RequiredHosts []string `json:"requiredHosts,omitempty" yaml:"requiredHosts,omitempty" mapstructure:"requiredHosts,omitempty"` } diff --git a/plugins/pdk/go/pdk/example_test.go b/plugins/pdk/go/pdk/example_test.go index 5678bddd4..bf98eba2f 100644 --- a/plugins/pdk/go/pdk/example_test.go +++ b/plugins/pdk/go/pdk/example_test.go @@ -8,6 +8,7 @@ package pdk_test import ( "testing" + "github.com/navidrome/navidrome/plugins/pdk/go/host" "github.com/navidrome/navidrome/plugins/pdk/go/pdk" "github.com/stretchr/testify/mock" ) @@ -136,48 +137,39 @@ func TestProcessJSONRequest(t *testing.T) { } // ============================================================================= -// Examples using stub types (Memory, HTTPRequest, HTTPResponse) +// HTTP requests go through the host HTTP service (host.HTTPSend), not +// pdk.NewHTTPRequest, which Navidrome does not enable. // ============================================================================= // FetchData demonstrates a plugin function that makes an HTTP request. func FetchData(url string) ([]byte, error) { - // Create and configure the HTTP request - // Note: SetHeader and SetBody work directly on the stub - no mocking needed! - req := pdk.NewHTTPRequest(pdk.MethodGet, url) - req.SetHeader("Accept", "application/json") - req.SetHeader("User-Agent", "MyPlugin/1.0") - - // Send the request - this is mocked because it requires host interaction - resp := req.Send() - - // Check status - works directly on the stub - if resp.Status() != 200 { - return nil, nil + resp, err := host.HTTPSend(host.HTTPRequest{ + Method: "GET", + URL: url, + Headers: map[string]string{ + "Accept": "application/json", + "User-Agent": "MyPlugin/1.0", + }, + }) + if err != nil { + return nil, err } - // Return body - works directly on the stub - return resp.Body(), nil + if resp.StatusCode != 200 { + return nil, nil + } + return resp.Body, nil } func TestFetchData(t *testing.T) { - pdk.ResetMock() + host.HTTPMock.ExpectedCalls = nil - // Create a stub response with test data expectedBody := []byte(`{"result": "success"}`) - stubResponse := pdk.NewStubHTTPResponse(200, map[string]string{ - "Content-Type": "application/json", - }, expectedBody) + host.HTTPMock.On("Send", mock.MatchedBy(func(req host.HTTPRequest) bool { + return req.Method == "GET" && req.URL == "https://api.example.com/data" && + req.Headers["Accept"] == "application/json" + })).Return(&host.HTTPResponse{StatusCode: 200, Body: expectedBody}, nil) - // Mock NewHTTPRequest to return a real HTTPRequest struct - // The struct methods (SetHeader, SetBody) work without mocking - pdk.PDKMock.On("NewHTTPRequest", pdk.MethodGet, "https://api.example.com/data"). - Return(&pdk.HTTPRequest{}) - - // Mock Send to return our stub response - pdk.PDKMock.On("Send", mock.AnythingOfType("*pdk.HTTPRequest")). - Return(stubResponse) - - // Call the function body, err := FetchData("https://api.example.com/data") if err != nil { @@ -188,21 +180,15 @@ func TestFetchData(t *testing.T) { t.Errorf("expected body %q, got %q", expectedBody, body) } - pdk.PDKMock.AssertExpectations(t) + host.HTTPMock.AssertExpectations(t) } func TestFetchData_NonOKStatus(t *testing.T) { - pdk.ResetMock() + host.HTTPMock.ExpectedCalls = nil - // Create a stub response with 404 status - stubResponse := pdk.NewStubHTTPResponse(404, nil, []byte("Not Found")) + host.HTTPMock.On("Send", mock.Anything). + Return(&host.HTTPResponse{StatusCode: 404, Body: []byte("Not Found")}, nil) - pdk.PDKMock.On("NewHTTPRequest", pdk.MethodGet, "https://api.example.com/missing"). - Return(&pdk.HTTPRequest{}) - pdk.PDKMock.On("Send", mock.AnythingOfType("*pdk.HTTPRequest")). - Return(stubResponse) - - // Call the function body, err := FetchData("https://api.example.com/missing") if err != nil { @@ -214,7 +200,7 @@ func TestFetchData_NonOKStatus(t *testing.T) { t.Errorf("expected nil body for 404, got %q", body) } - pdk.PDKMock.AssertExpectations(t) + host.HTTPMock.AssertExpectations(t) } // ProcessMemoryData demonstrates working with Memory type. @@ -292,23 +278,27 @@ func TestHTTPMethodString(t *testing.T) { // PostJSON demonstrates a more complex HTTP request with body. func PostJSON(url string, data []byte) (int, error) { - req := pdk.NewHTTPRequest(pdk.MethodPost, url) - req.SetHeader("Content-Type", "application/json") - req.SetBody(data) // Works directly on stub - - resp := req.Send() // This is mocked - return int(resp.Status()), nil + resp, err := host.HTTPSend(host.HTTPRequest{ + Method: "POST", + URL: url, + Headers: map[string]string{"Content-Type": "application/json"}, + Body: data, + }) + if err != nil { + return 0, err + } + return int(resp.StatusCode), nil } func TestPostJSON(t *testing.T) { - pdk.ResetMock() + host.HTTPMock.ExpectedCalls = nil - stubResponse := pdk.NewStubHTTPResponse(201, nil, nil) - - pdk.PDKMock.On("NewHTTPRequest", pdk.MethodPost, "https://api.example.com/items"). - Return(&pdk.HTTPRequest{}) - pdk.PDKMock.On("Send", mock.AnythingOfType("*pdk.HTTPRequest")). - Return(stubResponse) + host.HTTPMock.On("Send", host.HTTPRequest{ + Method: "POST", + URL: "https://api.example.com/items", + Headers: map[string]string{"Content-Type": "application/json"}, + Body: []byte(`{"name":"test"}`), + }).Return(&host.HTTPResponse{StatusCode: 201}, nil) status, err := PostJSON("https://api.example.com/items", []byte(`{"name":"test"}`)) @@ -320,5 +310,5 @@ func TestPostJSON(t *testing.T) { t.Errorf("expected status 201, got %d", status) } - pdk.PDKMock.AssertExpectations(t) + host.HTTPMock.AssertExpectations(t) } diff --git a/plugins/pdk/go/pdk/pdk.go b/plugins/pdk/go/pdk/pdk.go index 35394d700..456a238f5 100644 --- a/plugins/pdk/go/pdk/pdk.go +++ b/plugins/pdk/go/pdk/pdk.go @@ -111,6 +111,8 @@ func LogMemory(level LogLevel, m Memory) { } // NewHTTPRequest NewHTTPRequest returns a new `HTTPRequest`. +// +// Deprecated: Navidrome does not enable extism's http_request host function, so every request sent this way fails. Use host.HTTPSend instead. func NewHTTPRequest(method HTTPMethod, url string) *HTTPRequest { return extism.NewHTTPRequest(method, url) } diff --git a/plugins/pdk/go/pdk/pdk_stub.go b/plugins/pdk/go/pdk/pdk_stub.go index 8f2707d56..1a685eecc 100644 --- a/plugins/pdk/go/pdk/pdk_stub.go +++ b/plugins/pdk/go/pdk/pdk_stub.go @@ -114,6 +114,8 @@ func LogMemory(level LogLevel, m Memory) { } // NewHTTPRequest NewHTTPRequest returns a new `HTTPRequest`. +// +// Deprecated: Navidrome does not enable extism's http_request host function, so every request sent this way fails. Use host.HTTPSend instead. func NewHTTPRequest(method HTTPMethod, url string) *HTTPRequest { args := PDKMock.Called(method, url) var r0 *HTTPRequest diff --git a/plugins/pdk/rust/nd-pdk-host/src/nd_host_http.rs b/plugins/pdk/rust/nd-pdk-host/src/nd_host_http.rs index 1c44cd2f3..f1607112c 100644 --- a/plugins/pdk/rust/nd-pdk-host/src/nd_host_http.rs +++ b/plugins/pdk/rust/nd-pdk-host/src/nd_host_http.rs @@ -30,7 +30,7 @@ mod base64_bytes { } /// HTTPRequest represents an outbound HTTP request from a plugin. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct HTTPRequest { pub method: String, @@ -47,7 +47,7 @@ pub struct HTTPRequest { } /// HTTPResponse represents the response from an outbound HTTP request. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct HTTPResponse { pub status_code: i32, diff --git a/plugins/pdk/rust/nd-pdk-host/src/nd_host_library.rs b/plugins/pdk/rust/nd-pdk-host/src/nd_host_library.rs index b4b9b3fb0..966c1607f 100644 --- a/plugins/pdk/rust/nd-pdk-host/src/nd_host_library.rs +++ b/plugins/pdk/rust/nd-pdk-host/src/nd_host_library.rs @@ -7,7 +7,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; /// Library represents a music library with metadata. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct Library { pub id: i32, diff --git a/plugins/pdk/rust/nd-pdk-host/src/nd_host_matcher.rs b/plugins/pdk/rust/nd-pdk-host/src/nd_host_matcher.rs index be2819257..7fdf88b7f 100644 --- a/plugins/pdk/rust/nd-pdk-host/src/nd_host_matcher.rs +++ b/plugins/pdk/rust/nd-pdk-host/src/nd_host_matcher.rs @@ -7,7 +7,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; /// MatchOptions carries optional parameters for a match request. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct MatchOptions { #[serde(default)] diff --git a/plugins/pdk/rust/nd-pdk-host/src/nd_host_scrobbleretriever.rs b/plugins/pdk/rust/nd-pdk-host/src/nd_host_scrobbleretriever.rs index fe80d98f6..19010924a 100644 --- a/plugins/pdk/rust/nd-pdk-host/src/nd_host_scrobbleretriever.rs +++ b/plugins/pdk/rust/nd-pdk-host/src/nd_host_scrobbleretriever.rs @@ -7,7 +7,7 @@ use extism_pdk::*; use serde::{Deserialize, Serialize}; /// ScrobbleCountOptions carries optional parameters for counting user scrobbles -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct ScrobbleCountOptions { #[serde(default)] @@ -17,7 +17,7 @@ pub struct ScrobbleCountOptions { } /// ScrobbleOptions carries optional parameters for retrieving user scrobbles -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct ScrobbleOptions { #[serde(default)] @@ -31,7 +31,7 @@ pub struct ScrobbleOptions { } /// ScrobbleRef represents one instance of a scrobble (instance id, file id, submission time) -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct ScrobbleRef { pub id: i64, diff --git a/plugins/pdk/rust/nd-pdk-host/src/nd_host_task.rs b/plugins/pdk/rust/nd-pdk-host/src/nd_host_task.rs index 4f43e165c..ac82fdd08 100644 --- a/plugins/pdk/rust/nd-pdk-host/src/nd_host_task.rs +++ b/plugins/pdk/rust/nd-pdk-host/src/nd_host_task.rs @@ -30,7 +30,7 @@ mod base64_bytes { } /// QueueConfig holds configuration for a task queue. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct QueueConfig { pub concurrency: i32, @@ -41,7 +41,7 @@ pub struct QueueConfig { } /// TaskInfo holds the current state of a task. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct TaskInfo { pub status: String, diff --git a/plugins/pdk/rust/nd-pdk-host/src/nd_host_users.rs b/plugins/pdk/rust/nd-pdk-host/src/nd_host_users.rs index faa795bb9..b31b39fb4 100644 --- a/plugins/pdk/rust/nd-pdk-host/src/nd_host_users.rs +++ b/plugins/pdk/rust/nd-pdk-host/src/nd_host_users.rs @@ -8,7 +8,7 @@ use serde::{Deserialize, Serialize}; /// User represents a Navidrome user with minimal information exposed to plugins. /// Sensitive fields like password, email, and internal IDs are intentionally excluded. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct User { pub user_name: String, 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/release/goreleaser.yml b/release/goreleaser.yml index 103f2beaf..e7610f25e 100644 --- a/release/goreleaser.yml +++ b/release/goreleaser.yml @@ -114,7 +114,7 @@ release: ## Where to go next? * Read installation instructions on our [website](https://www.navidrome.org/docs/installation/). - * Host Navidrome on [PikaPods](https://www.pikapods.com/pods/navidrome) or [Danian](https://danian.co/navidrome?nd) for a simple cloud solution. + * Host Navidrome on [PikaPods](https://www.pikapods.com/pods?run=navidrome) or [Zenith](https://zenith.hosting/host/navidrome?ref=navidrome) for a simple cloud solution. * Reach out on [Discord](https://discord.gg/xh7j7yF), [Reddit](https://www.reddit.com/r/navidrome/) and [Twitter](https://twitter.com/navidrome)! # Add the MSI installers to the release diff --git a/release/linux/postinstall.sh b/release/linux/postinstall.sh index ed39fb127..007114ef5 100644 --- a/release/linux/postinstall.sh +++ b/release/linux/postinstall.sh @@ -13,15 +13,16 @@ if [ ! -f /etc/navidrome/navidrome.toml ]; then printf "MusicFolder = \"/opt/navidrome/music\"\n" >> /etc/navidrome/navidrome.toml fi +# Older versions created these folders as root when this script ran `navidrome`. They were created empty, +# so no -R. Real dirs only, never following links: navidrome owns /var/lib/navidrome and could plant symlinks. +find /var/lib/navidrome/cache /var/lib/navidrome/artwork /var/lib/navidrome/plugins -maxdepth 0 -type d -user root -exec chown -h navidrome:navidrome {} + 2>/dev/null + postinstall_flag="/var/lib/navidrome/.installed" if [ ! -f "$postinstall_flag" ]; then # The primary reason why this would fail is if the service was already installed AND # someone manually removed the .installed flag. In this case, ignore the error navidrome service install --user navidrome --working-directory /var/lib/navidrome --configfile /etc/navidrome/navidrome.toml || : - # Any `navidrome` command will make a cache. Make sure that this is properly owned by the Navidrome user - # and not by root - chown navidrome:navidrome /var/lib/navidrome/cache touch "$postinstall_flag" else navidrome service stop --configfile /etc/navidrome/navidrome.toml && navidrome service start --configfile /etc/navidrome/navidrome.toml diff --git a/release/wix/SettingsDlg.wxs b/release/wix/SettingsDlg.wxs index 4a83f91da..cdc625434 100644 --- a/release/wix/SettingsDlg.wxs +++ b/release/wix/SettingsDlg.wxs @@ -16,7 +16,7 @@ - + diff --git a/release/wix/navidrome.wxs b/release/wix/navidrome.wxs index 6d94bab9d..71a6de6b5 100644 --- a/release/wix/navidrome.wxs +++ b/release/wix/navidrome.wxs @@ -14,6 +14,7 @@ + @@ -23,55 +24,65 @@ - - - + + + + + + + + + + + + - + - - + + + + - - - + + + - - - + + + + + - - - - - - + + + + + - - - - - + + + + + - - - + + + - - - - - + + + @@ -81,7 +92,12 @@ + + NOT MSI_INSTALLATIONDIRECTORY + + + NOT MSI_INSTALLATIONDIRECTORY Not Installed AND NOT WIX_UPGRADE_DETECTED diff --git a/resources/i18n/de.json b/resources/i18n/de.json index 1a516d393..400a2364d 100644 --- a/resources/i18n/de.json +++ b/resources/i18n/de.json @@ -93,7 +93,8 @@ "addToPlaylist": "Zu einer Wiedergabeliste hinzufügen", "download": "Herunterladen", "info": "Mehr Informationen", - "share": "Freigabe erstellen" + "share": "Freigabe erstellen", + "refresh": "Metadaten aktualisieren" }, "lists": { "all": "Alle", @@ -155,11 +156,13 @@ "newPassword": "Neues Passwort", "token": "Token", "lastAccessAt": "Letzter Zugriff am", - "libraries": "Bibliotheken" + "libraries": "Bibliotheken", + "scrobbleFilter": "Scrobble-Filter" }, "helperTexts": { "name": "Die Änderung wird erst nach dem nächsten Login gültig", - "libraries": "Wähle spezifische Bibliotheken für diesen Benutzer, oder leer lassen für Standard Bibliotheken" + "libraries": "Wähle spezifische Bibliotheken für diesen Benutzer, oder leer lassen für Standard Bibliotheken", + "scrobbleFilter": "Titel, die diesen Regeln für intelligente Wiedergabelisten entsprechen, werden nicht an Last.fm, ListenBrainz oder Scrobbler-Plugins übermittelt. Es werden dieselbe JSON-Syntax und dasselbe Verhalten wie bei intelligenten Wiedergabelisten verwendet. Beispiel: {\"all\":[{\"lt\":{\"rating\":4}}]}. Ist das Feld leer, wird alles gescrobbelt. Lokale Wiedergabezahlen bleiben davon unberührt." }, "notifications": { "created": "Benutzer erstellt", @@ -173,7 +176,8 @@ "adminAutoLibraries": "Administrator-Benutzer haben automatisch Zugriff auf alle Bibliotheken" }, "validation": { - "librariesRequired": "Mindestens eine Bibliothek muss für nicht-administrator Benutzer ausgewählt sein" + "librariesRequired": "Mindestens eine Bibliothek muss für nicht-administrator Benutzer ausgewählt sein", + "invalidScrobbleFilter": "Es müssen gültige Regeln für intelligente Wiedergabelisten sein. Limit, Offset und Aktualisierungsverzögerung werden nicht unterstützt." } }, "player": { @@ -196,6 +200,9 @@ "targetFormat": "Zielformat", "defaultBitRate": "Standardbitrate", "command": "Befehl" + }, + "choices": { + "noDefaultBitRate": "Keine" } }, "playlist": { @@ -210,7 +217,8 @@ "songCount": "Titelanzahl", "comment": "Kommentar", "sync": "Auto-Import", - "path": "Importieren aus" + "path": "Importieren aus", + "starred": "Favorit" }, "actions": { "selectPlaylist": "Wiedergabeliste auswählen:", @@ -390,6 +398,7 @@ "invalidJson": "Konfiguration muss valides JSON sein" }, "messages": { + "idHelp": "Die Plugin-ID, abgeleitet vom Dateinamen. Verwende diese, um in Konfigurationen (z. B. bei Agents) auf dieses Plugin zu verweisen.", "configHelp": "Plugin mit Schlüssel-Werte Paaren konfigurieren. Leer lassen wenn das Plugin keine Konfiguration benötigt.", "clickPermissions": "Berechtigung anklicken für mehr Details", "noConfig": "Keine Konfiguration gesetzt", @@ -564,7 +573,7 @@ "noPlaylistsAvailable": "Keine Wiedergabeliste verfügbar", "delete_user_title": "Benutzer '%{name}' löschen", "delete_user_content": "Möchtest du diesen Benutzer und alle seine Daten (einschließlich Wiedergabelisten und Einstellungen) wirklich löschen?", - "notifications_blocked": "Sie haben Benachrichtigungen für diese Seite in den Einstellungen Ihres Browsers blockiert", + "notifications_blocked": "Benachrichtigungen für diese Seite sind in den Browsereinstellungen blockiert", "notifications_not_available": "Dieser Browser unterstützt keine Desktop-Benachrichtigungen", "lastfmLinkSuccess": "Last.fm Verbindung hergestellt und scrobbling aktiviert", "lastfmLinkFailure": "Last.fm konnte nicht verbunden werden", @@ -599,7 +608,8 @@ "coverUploaded": "Cover aktualisiert", "coverRemoved": "Cover entfernt", "coverUploadError": "Fehler beim Hochladen des Covers", - "coverRemoveError": "Fehler beim Entfernen des Covers" + "coverRemoveError": "Fehler beim Entfernen des Covers", + "metadataRefreshStarted": "Aktualisiert Metadaten im Hintergrund" }, "menu": { "library": "Bibliothek", @@ -634,7 +644,8 @@ "multipleLibraries": "%{selected} von %{total} Bibliotheken", "selectLibraries": "Bibliotheken auswählen", "none": "Keine" - } + }, + "onlyFavourites": "Nur Favoriten anzeigen" }, "player": { "playListsText": "Warteschlange abspielen", @@ -720,4 +731,4 @@ "empty": "Keine Wiedergabe", "minutesAgo": "Vor %{smart_count} Minute |||| Vor %{smart_count} Minuten" } -} \ No newline at end of file +} diff --git a/resources/i18n/el.json b/resources/i18n/el.json index 3876c2e33..37882ea99 100644 --- a/resources/i18n/el.json +++ b/resources/i18n/el.json @@ -38,7 +38,9 @@ "missing": "Απών", "libraryName": "Βιβλιοθήκη", "composer": "Συνθέτης", - "disc": "Δίσκος %{discNumber}" + "disc": "Δίσκος %{discNumber}", + "albumGain": "Κέρδος άλμπουμ", + "trackGain": "Κέρδος παρακολούθησης" }, "actions": { "addToQueue": "Αναπαραγωγη Μετα", @@ -91,7 +93,8 @@ "addToPlaylist": "Προσθηκη στη λιστα αναπαραγωγης", "download": "Ληψη", "info": "Εμφάνιση Πληροφοριών", - "share": "Μερίδιο" + "share": "Μερίδιο", + "refresh": "Ανανέωση μεταδεδομένων" }, "lists": { "all": "Όλα", @@ -153,11 +156,13 @@ "newPassword": "Νέος Κωδικός Πρόσβασης", "token": "Token", "lastAccessAt": "Τελευταία Πρόσβαση", - "libraries": "Βιβλιοθήκες" + "libraries": "Βιβλιοθήκες", + "scrobbleFilter": "Φίλτρο Scrobble" }, "helperTexts": { "name": "Αλλαγές στο όνομα σας θα εφαρμοστούν στην επόμενη σύνδεση", - "libraries": "Επιλέξτε συγκεκριμένες βιβλιοθήκες για αυτόν τον χρήστη, ή αφήστε την κενή για να χρησιμοποιήσετε την προεπιλεγμένη βιβλιοθήκη" + "libraries": "Επιλέξτε συγκεκριμένες βιβλιοθήκες για αυτόν τον χρήστη, ή αφήστε την κενή για να χρησιμοποιήσετε την προεπιλεγμένη βιβλιοθήκη", + "scrobbleFilter": "Τα τραγούδια που αντιστοιχούν σε αυτούς τους κανόνες έξυπνης λίστας αναπαραγωγής δεν αποστέλλονται στα πρόσθετα Last.fm, ListenBrainz ή scrobbler. Χρησιμοποιεί την ίδια σύνταξη και συμπεριφορά JSON με τις έξυπνες λίστες αναπαραγωγής. Παράδειγμα: {\"all\":[{\"lt\":{\"rating\":4}}]}." }, "notifications": { "created": "Ο χρήστης δημιουργήθηκε", @@ -171,7 +176,8 @@ "adminAutoLibraries": "Οι χρήστες διαχειριστές έχουν αυτόματα πρόσβαση σε όλες τις βιβλιοθήκες" }, "validation": { - "librariesRequired": "Πρέπει να επιλεγεί τουλάχιστον μία βιβλιοθήκη για χρήστες που δεν είναι διαχειριστές" + "librariesRequired": "Πρέπει να επιλεγεί τουλάχιστον μία βιβλιοθήκη για χρήστες που δεν είναι διαχειριστές", + "invalidScrobbleFilter": "Πρέπει να υπάρχουν έγκυροι κανόνες έξυπνης λίστας αναπαραγωγής. Δεν υποστηρίζονται το όριο, η μετατόπιση και η καθυστέρηση ανανέωσης." } }, "player": { @@ -194,6 +200,9 @@ "targetFormat": "Μορφη Προορισμου", "defaultBitRate": "Προκαθορισμένος Ρυθμός Bit", "command": "Εντολή" + }, + "choices": { + "noDefaultBitRate": "" } }, "playlist": { @@ -208,7 +217,8 @@ "songCount": "Τραγούδια", "comment": "Σχόλιο", "sync": "Αυτόματη εισαγωγή", - "path": "Εισαγωγή από" + "path": "Εισαγωγή από", + "starred": "Ευνοούμενος" }, "actions": { "selectPlaylist": "Επιλέξτε μια λίστα αναπαραγωγής:", @@ -401,7 +411,8 @@ "requiredHosts": "Απαιτούμενοι hosts", "configValidationError": "Η επικύρωση διαμόρφωσης απέτυχε:", "schemaRenderError": "Δεν είναι δυνατή η απόδοση της φόρμας διαμόρφωσης. Το σχήμα της προσθήκης ενδέχεται να μην είναι έγκυρο.", - "allowWriteAccessHelp": "Όταν είναι ενεργοποιημένο, το πρόσθετο μπορεί να τροποποιήσει αρχεία στους καταλόγους της βιβλιοθήκης. Από προεπιλογή, τα πρόσθετα έχουν πρόσβαση μόνο για ανάγνωση." + "allowWriteAccessHelp": "Όταν είναι ενεργοποιημένο, το πρόσθετο μπορεί να τροποποιήσει αρχεία στους καταλόγους της βιβλιοθήκης. Από προεπιλογή, τα πρόσθετα έχουν πρόσβαση μόνο για ανάγνωση.", + "idHelp": "" }, "placeholders": { "configKey": "κλειδί", @@ -597,7 +608,8 @@ "coverUploaded": "Το εξώφυλλο ενημερώθηκε", "coverRemoved": "Το εξώφυλλο αφαιρέθηκε", "coverUploadError": "Σφάλμα κατά τη μεταφόρτωση του εξωφύλλου", - "coverRemoveError": "Σφάλμα κατά την αφαίρεση του εξωφύλλου" + "coverRemoveError": "Σφάλμα κατά την αφαίρεση του εξωφύλλου", + "metadataRefreshStarted": "Ανανέωση μεταδεδομένων στο παρασκήνιο" }, "menu": { "library": "Βιβλιοθήκη", @@ -632,7 +644,8 @@ "multipleLibraries": "%{selected} από %{total} Βιβλιοθήκες", "selectLibraries": "Επιλέξτε βιβλιοθήκες", "none": "Κανένα" - } + }, + "onlyFavourites": "Εμφάνιση μόνο αγαπημένων" }, "player": { "playListsText": "Ουρά Αναπαραγωγής", diff --git a/resources/i18n/fi.json b/resources/i18n/fi.json index 0e6149f87..f79d71f94 100644 --- a/resources/i18n/fi.json +++ b/resources/i18n/fi.json @@ -93,7 +93,8 @@ "addToPlaylist": "Lisää soittolistaan", "download": "Lataa", "info": "Info", - "share": "Jaa" + "share": "Jaa", + "refresh": "Päivitä metatiedot" }, "lists": { "all": "Kaikki", @@ -155,11 +156,13 @@ "newPassword": "Uusi salasana", "token": "Avain", "lastAccessAt": "Viimeisin käyttö", - "libraries": "Kirjastot" + "libraries": "Kirjastot", + "scrobbleFilter": "Scrobble-suodatin" }, "helperTexts": { "name": "Nimen muutos tulee voimaan kun seuraavan kerran kirjaudut sisään", - "libraries": "Valitse tietyt kirjastot tälle käyttäjälle tai jätä tyhjäksi käyttääksesi oletuskirjastoja" + "libraries": "Valitse tietyt kirjastot tälle käyttäjälle tai jätä tyhjäksi käyttääksesi oletuskirjastoja", + "scrobbleFilter": "Älykkään soittolistan sääntöihin osumat kappaleet ohitetaan Last.fm-, ListenBrainz- ja skrobbausliitännäisissä. Käyttää samaa JSON-syntaksia ja toimintalogiikkaa kuin älykkäät soittolistat. Esimerkki: {\"all\":[{\"lt\":{\"rating\":4}}]}. Jätä tyhjäksi, jos haluat skrobata kaiken. Paikallisiin toistokertoihin tämä ei vaikuta." }, "notifications": { "created": "Käyttäjä luotu", @@ -173,7 +176,8 @@ "adminAutoLibraries": "Admin-käyttäjillä on automaattisesti pääsy kaikkiin kirjastoihin" }, "validation": { - "librariesRequired": "Vähintään yksi kirjasto on valittava ei-admin käyttäjille" + "librariesRequired": "Vähintään yksi kirjasto on valittava ei-admin käyttäjille", + "invalidScrobbleFilter": "Sääntöjen pitää olla kelvollisia älykkään soittolistan sääntöjä. Rajaus (limit), siirtymä (offset) ja päivitysviive eivät toimi tässä." } }, "player": { @@ -196,6 +200,9 @@ "targetFormat": "Kohde formaatti", "defaultBitRate": "Oletus bittinopeus", "command": "Komento" + }, + "choices": { + "noDefaultBitRate": "Ei mikään" } }, "playlist": { @@ -210,7 +217,8 @@ "songCount": "Kappaleita", "comment": "Kommentti", "sync": "Automaattinen tuonti", - "path": "Tuo" + "path": "Tuo", + "starred": "Suosikki" }, "actions": { "selectPlaylist": "Valitse soittolista:", @@ -403,7 +411,8 @@ "requiredHosts": "Vaaditut palvelimet", "configValidationError": "Määrityksen validointi epäonnistui:", "schemaRenderError": "Konfiguraatiolomaketta ei voi näyttää. Lisäosan skeema saattaa olla virheellinen.", - "allowWriteAccessHelp": "Kun otettu käyttöön, liitännäinen voi muokata tiedostoja kirjastohakemistoissa. Oletuksena liitännäisillä on vain luku -oikeus." + "allowWriteAccessHelp": "Kun otettu käyttöön, liitännäinen voi muokata tiedostoja kirjastohakemistoissa. Oletuksena liitännäisillä on vain luku -oikeus.", + "idHelp": "Lisäosan ID, joka johdetaan sen tiedostonimestä. Käytä sitä viitatessasi tähän lisäosaan määritelmäasetuksissa, kuten agenteissa." }, "placeholders": { "configKey": "avain", @@ -599,7 +608,11 @@ "coverUploaded": "Kansikuva päivitetty", "coverRemoved": "Kansikuva poistettu", "coverUploadError": "Virhe ladattaessa kansikuvaa", - "coverRemoveError": "Virhe poistettaessa kansikuvaa" + "coverRemoveError": "Virhe poistettaessa kansikuvaa", + "metadataRefreshStarted": "Metatietoja päivitetään taustalla", + "quickConnectApproved": "", + "quickConnectInvalidCode": "", + "quickConnectError": "" }, "menu": { "library": "Kirjasto", @@ -634,6 +647,15 @@ "multipleLibraries": "%{selected} / %{total} kirjastoa", "selectLibraries": "Valitse kirjastot", "none": "Ei mitään" + }, + "onlyFavourites": "Näytä vain suosikit", + "quickConnect": { + "name": "", + "code": "", + "help": "", + "confirm": "", + "continue": "", + "approve": "" } }, "player": { diff --git a/resources/i18n/gl.json b/resources/i18n/gl.json index 444998d03..47f2f4272 100644 --- a/resources/i18n/gl.json +++ b/resources/i18n/gl.json @@ -93,7 +93,8 @@ "addToPlaylist": "Engadir a Lista", "download": "Descargar", "info": "Obter info", - "share": "Compartir" + "share": "Compartir", + "refresh": "Actualizar metadatos" }, "lists": { "all": "Todo", @@ -155,11 +156,13 @@ "newPassword": "Novo contrasinal", "token": "Token", "lastAccessAt": "Último acceso", - "libraries": "Bibliotecas" + "libraries": "Bibliotecas", + "scrobbleFilter": "Filtro para scrobble" }, "helperTexts": { "name": "Os cambios no nome aplicaranse a próxima vez que accedas", - "libraries": "Selecciona bibliotecas específicas para esta usuaria, ou deixa baleiro para usar as bibliotecas por defecto" + "libraries": "Selecciona bibliotecas específicas para esta usuaria, ou deixa baleiro para usar as bibliotecas por defecto", + "scrobbleFilter": "As cancións que concorden coas regras desta lista intelixente non se envían a Last.fm, ListenBrainz ou complementos similares. O filtro usa a mesma sintaxe JSON e comportamento que as listas de reprodución intelixentes. Exemplo: {\"all\":[{\"lt\":{\"rating\":4}}]}. Deixar baleiro para enviar todo. Non lle afecta ao número de reproducións locais." }, "notifications": { "created": "Creouse a usuaria", @@ -173,7 +176,8 @@ "adminAutoLibraries": "As usuarias Admin teñen acceso por defecto a todas as bibliotecas" }, "validation": { - "librariesRequired": "Debes seleccionar polo menos unha biblioteca para usuarias non admins" + "librariesRequired": "Debes seleccionar polo menos unha biblioteca para usuarias non admins", + "invalidScrobbleFilter": "" } }, "player": { @@ -196,6 +200,9 @@ "targetFormat": "Formato de destino", "defaultBitRate": "Taxa de bit por defecto", "command": "Orde" + }, + "choices": { + "noDefaultBitRate": "" } }, "playlist": { @@ -210,7 +217,8 @@ "songCount": "Cancións", "comment": "Comentario", "sync": "Autoimportación", - "path": "Importar desde" + "path": "Importar desde", + "starred": "Favorita" }, "actions": { "selectPlaylist": "Elixe unha lista:", @@ -403,7 +411,8 @@ "requiredHosts": "Servidores requeridos", "configValidationError": "Fallou a comprobación da configuración:", "schemaRenderError": "Non se puido aplicar a configuración. O esquema do complemento podería non ser válido.", - "allowWriteAccessHelp": "A activalo, este complemento pode modificar ficheiros nos directorios da biblioteca. Por defecto os complementos teñen acceso de só-lectura." + "allowWriteAccessHelp": "A activalo, este complemento pode modificar ficheiros nos directorios da biblioteca. Por defecto os complementos teñen acceso de só-lectura.", + "idHelp": "O ID do complemento, derivado do seu nome de ficheiro. Utilízao cando te refiras ao complemento nas opcións de configuración, como en Agents." }, "placeholders": { "configKey": "clave", @@ -435,7 +444,7 @@ "minValue": "Ten que ter polo menos %{min}", "maxValue": "Ten que ter %{max} ou menos", "number": "Ten que ser un número", - "email": "Ten que ser un email válido", + "email": "Ten que ser un correo válido", "oneOf": "Ten que ser un de: %{options}", "regex": "Ten que ter un formato específico (regexp): %{pattern}", "unique": "Ten que ser único", @@ -509,7 +518,7 @@ } }, "message": { - "about": "Acerca de", + "about": "Sobre", "are_you_sure": "Tes certeza?", "bulk_delete_content": "Tes a certeza de querer borrar a %{name} |||| Tes a certeza de querer eleminar estes %{smart_count} elementos?", "bulk_delete_title": "Eliminar %{name} |||| Eliminar %{smart_count} %{name}", @@ -599,7 +608,8 @@ "coverUploaded": "Subiuse a capa", "coverRemoved": "Retirouse a capa", "coverUploadError": "Erro ao subir a capa", - "coverRemoveError": "Erro ao retirar a capa" + "coverRemoveError": "Erro ao retirar a capa", + "metadataRefreshStarted": "Actualizar en segundo plano os metadatos" }, "menu": { "library": "Biblioteca", @@ -626,7 +636,7 @@ } }, "albumList": "Álbums", - "about": "Acerca de", + "about": "Sobre", "playlists": "Listas de reprodución", "sharedPlaylists": "Listas compartidas", "librarySelector": { @@ -634,7 +644,8 @@ "multipleLibraries": "%{selected} de %{total} Bibliotecas", "selectLibraries": "Seleccionar Bibliotecas", "none": "Ningunha" - } + }, + "onlyFavourites": "Mostrar só as favoritas" }, "player": { "playListsText": "Reproducir cola", diff --git a/resources/i18n/nl.json b/resources/i18n/nl.json index 46c3df9de..cb276b8f4 100644 --- a/resources/i18n/nl.json +++ b/resources/i18n/nl.json @@ -93,7 +93,8 @@ "addToPlaylist": "Toevoegen aan afspeellijst", "download": "Downloaden", "info": "Meer info", - "share": "Delen" + "share": "Delen", + "refresh": "Metadata verversen" }, "lists": { "all": "Alle", @@ -155,11 +156,13 @@ "newPassword": "Nieuw wachtwoord", "token": "Token", "lastAccessAt": "Meest recente toegang", - "libraries": "Bibliotheken" + "libraries": "Bibliotheken", + "scrobbleFilter": "Scrobble filter" }, "helperTexts": { "name": "Naamswijziging wordt pas zichtbaar bij de volgende login", - "libraries": "Selecteer specifieke bibliotheken voor deze gebruiker, of laat leeg om de standaardbiblliotheken te gebruiken" + "libraries": "Selecteer specifieke bibliotheken voor deze gebruiker, of laat leeg om de standaardbiblliotheken te gebruiken", + "scrobbleFilter": "Nummers binnen deze slimme afspeellijst regels worden niet gestuurd naar Last.fm, ListenBrainz of scrobbler plugins. Gebruikt dezelfde JSON syntaxis en gedrag als slimme afspeellijsten. Voorbeeld: {\"all\":[{\"lt\":{\"rating\":4}}]}. Leeglaten om alles te scrobblen. Lokale afspeelaantallen vallen hierbuiten." }, "notifications": { "created": "Aangemaakt door gebruiker", @@ -173,7 +176,8 @@ "adminAutoLibraries": "Admin gebruikers hebben automatisch toegang tot alle bibliotheken" }, "validation": { - "librariesRequired": "Minstens één bibliotheek moet geselecteerd worden voor niet-admin gebruikers" + "librariesRequired": "Minstens één bibliotheek moet geselecteerd worden voor niet-admin gebruikers", + "invalidScrobbleFilter": "Moeten geldige slimme afspeellijst regels zijn. Geen ondersteuning voor Limit, Offset en Refresh Delay" } }, "player": { @@ -196,6 +200,9 @@ "targetFormat": "Doelformaat", "defaultBitRate": "Standaard bitrate", "command": "Commando" + }, + "choices": { + "noDefaultBitRate": "Geen" } }, "playlist": { @@ -210,7 +217,8 @@ "songCount": "Nummers", "comment": "Commentaar", "sync": "Auto-importeren", - "path": "Importeer vanuit" + "path": "Importeer vanuit", + "starred": "Favoriet" }, "actions": { "selectPlaylist": "Selecteer een afspeellijst:", @@ -403,7 +411,8 @@ "requiredHosts": "Benodigde hosts", "configValidationError": "Configuratiecheck mislukt", "schemaRenderError": "Kan het configuratieformulier niet verwerken. Het plugin schema is wellicht ongeldig.", - "allowWriteAccessHelp": "Met dit ingeschakeld, kan de plug-in bestanden bewerken in de bibliotheekmappen. Standaard kunnen plug-ins alleen lezen." + "allowWriteAccessHelp": "Met dit ingeschakeld, kan de plug-in bestanden bewerken in de bibliotheekmappen. Standaard kunnen plug-ins alleen lezen.", + "idHelp": "De plugin ID, afgeleid van zijn bestandsnaam. Gebruik dit als gerefereerd wordt aan deze plugin in de configuratie opties, zoals Agenten." }, "placeholders": { "configKey": "Sleutel", @@ -599,7 +608,11 @@ "coverUploaded": "Albumhoes bijgewerkt", "coverRemoved": "Albumhoes verwijderd", "coverUploadError": "Fout bij het toevoegen albumhoes", - "coverRemoveError": "Fout bij verwijderen albumhoes" + "coverRemoveError": "Fout bij verwijderen albumhoes", + "metadataRefreshStarted": "Metadata verversen op de achtergrond", + "quickConnectApproved": "", + "quickConnectInvalidCode": "", + "quickConnectError": "" }, "menu": { "library": "Bibliotheek", @@ -634,6 +647,15 @@ "multipleLibraries": "%{selected} van %{total} bibliotheken", "selectLibraries": "Selecteer bibliotheken", "none": "Geen" + }, + "onlyFavourites": "Toon alleen favorieten", + "quickConnect": { + "name": "", + "code": "", + "help": "", + "confirm": "", + "continue": "", + "approve": "" } }, "player": { diff --git a/resources/i18n/pl.json b/resources/i18n/pl.json index 6229798e9..ff697d64c 100644 --- a/resources/i18n/pl.json +++ b/resources/i18n/pl.json @@ -37,7 +37,10 @@ "sampleRate": "Częstotliwość próbkowania", "missing": "Brak", "libraryName": "Biblioteka", - "composer": "Kompozytor" + "composer": "Kompozytor", + "disc": "Dysk %{discNumber}", + "albumGain": "Wzmocnienie albumu", + "trackGain": "Wzmocnienie utworu" }, "actions": { "addToQueue": "Odtwarzaj Później", @@ -90,7 +93,8 @@ "addToPlaylist": "Dodaj do Playlisty", "download": "Pobierz", "info": "Zdobądź Informacje", - "share": "Udostępnij" + "share": "Udostępnij", + "refresh": "Odśwież Metadane" }, "lists": { "all": "Wszystkie", @@ -152,11 +156,13 @@ "newPassword": "Nowe hasło", "token": "Token", "lastAccessAt": "Ostatnia Aktywność", - "libraries": "Biblioteki" + "libraries": "Biblioteki", + "scrobbleFilter": "Filtr scrobblowania" }, "helperTexts": { "name": "Zmiana nazwy będzie widoczna przy następnym logowaniu", - "libraries": "Wybierz biblioteki dla użytkownika lub pozostaw pustę, aby użyć domyślnej biblioteki" + "libraries": "Wybierz biblioteki dla użytkownika lub pozostaw pustę, aby użyć domyślnej biblioteki", + "scrobbleFilter": "Utwory spełniające kryteria tych inteligentnych list odtwarzania nie są wysyłane do serwisów Last.fm, ListenBrainz ani do wtyczek typu scrobbler. Wykorzystywana jest tu ta sama składnia JSON i zasada działania, co w przypadku inteligentnych list odtwarzania. Przykład: {\"all\":[{\"lt\":{\"rating\":4}}]}. Pozostawienie pola pustego spowoduje scrobblowanie wszystkich utworów. Nie wpływa to na lokalne liczniki odtworzeń." }, "notifications": { "created": "Dodano użytkownika", @@ -170,7 +176,8 @@ "adminAutoLibraries": "Administratorzy automatycznie mają dostęp do wszystkich bibliotek" }, "validation": { - "librariesRequired": "Przynajmniej jedna biblioteka musi być wybrana dla zwykłego użytkownika" + "librariesRequired": "Przynajmniej jedna biblioteka musi być wybrana dla zwykłego użytkownika", + "invalidScrobbleFilter": "Reguły inteligentnej listy odtwarzania muszą być poprawne. Parametry limitu, przesunięcia oraz opóźnienia odświeżania nie są obsługiwane." } }, "player": { @@ -193,6 +200,9 @@ "targetFormat": "Format Docelowy", "defaultBitRate": "Domyślny Bit Rate", "command": "Komenda" + }, + "choices": { + "noDefaultBitRate": "" } }, "playlist": { @@ -207,7 +217,8 @@ "songCount": "Liczba utworów", "comment": "Komentarz", "sync": "Import automatyczny", - "path": "Zaimportuj z" + "path": "Zaimportuj z", + "starred": "Ulubione" }, "actions": { "selectPlaylist": "Wybierz playlistę:", @@ -353,7 +364,8 @@ "allUsers": "Zezwalaj wszystkim użytkownikom", "selectedUsers": "Wybrani użytkownicy", "allLibraries": "Zezwalaj dla wszystkich bibliotek", - "selectedLibraries": "Wybrane biblioteki" + "selectedLibraries": "Wybrane biblioteki", + "allowWriteAccess": "Zezwól na zapis" }, "sections": { "status": "Status", @@ -398,7 +410,9 @@ "librariesRequired": "Wtyczka wymaga dostępu do informacji o bibliotece. Wybierz, dla której biblioteki zezwolić dostęp, lub włącz 'Zezwalaj dla wszystkich bibliotek'.", "requiredHosts": "Wymagane hosty", "configValidationError": "Weryfikacja konfiguracji nie powiodła się:", - "schemaRenderError": "Nie można wyrenderować formularza konfiguracji. Schemat wtyczki może być nieprawidłowy." + "schemaRenderError": "Nie można wyrenderować formularza konfiguracji. Schemat wtyczki może być nieprawidłowy.", + "allowWriteAccessHelp": "Po włączeniu wtyczka może modyfikować pliki w katalogach bibliotek. Domyślnie wtyczki mają dostęp tylko do odczytu.", + "idHelp": "" }, "placeholders": { "configKey": "klucz", @@ -588,7 +602,14 @@ "remove_all_missing_content": "Czy chcesz usunąć wszystkie brakujące pliki z bazy danych? Spowoduje to trwałe usunięcie wszelkich odniesień do tych plików, takich jak liczba odtworzeń, czy oceny.", "noSimilarSongsFound": "Brak podobnych utworów", "noTopSongsFound": "Brak najlepszych utworów", - "startingInstantMix": "Ładowanie Natychmiastowego Miksu..." + "startingInstantMix": "Ładowanie Natychmiastowego Miksu...", + "uploadCover": "Prześlij Okładkę", + "removeCover": "Usuń Okładkę", + "coverUploaded": "Okładka zaktualizowana", + "coverRemoved": "Okładka usunięta", + "coverUploadError": "Błąd przesyłania okładki", + "coverRemoveError": "Błąd usuwania okładki", + "metadataRefreshStarted": "Odświeżanie metadanych w tle" }, "menu": { "library": "Biblioteka", @@ -623,7 +644,8 @@ "multipleLibraries": "%{selected} z %{total} Bibliotek", "selectLibraries": "Wybierz Biblioteki", "none": "Żadna" - } + }, + "onlyFavourites": "Pokaż tylko ulubione" }, "player": { "playListsText": "Kolejka Odtwarzania", @@ -674,7 +696,8 @@ "exportSuccess": "Konfiguracja wyeksportowana do schowka w formacie TOML", "exportFailed": "Błąd kopiowania konfiguracji", "devFlagsHeader": "Flagi Rozwojowe (mogą ulec zmianie/usunięciu)", - "devFlagsComment": "To są ustawienia eksperymentalne i mogą zostać usunięte w przyszłych wydaniach" + "devFlagsComment": "To są ustawienia eksperymentalne i mogą zostać usunięte w przyszłych wydaniach", + "downloadToml": "Konfiguracja Pobierania (TOML)" } }, "activity": { diff --git a/resources/i18n/pt-br.json b/resources/i18n/pt-br.json index ccc5f872b..238919259 100644 --- a/resources/i18n/pt-br.json +++ b/resources/i18n/pt-br.json @@ -35,12 +35,12 @@ "rawTags": "Tags originais", "bitDepth": "Profundidade de bits", "sampleRate": "Taxa de amostragem", - "albumGain": "Ganho do álbum", - "trackGain": "Ganho da faixa", "missing": "Ausente", "libraryName": "Biblioteca", "composer": "Compositor", - "disc": "Disco %{discNumber}" + "disc": "Disco %{discNumber}", + "albumGain": "Ganho do álbum", + "trackGain": "Ganho da faixa" }, "actions": { "addToQueue": "Adicionar à fila", @@ -93,8 +93,8 @@ "addToPlaylist": "Adicionar à playlist", "download": "Baixar", "info": "Detalhes", - "refresh": "Atualizar Metadados", - "share": "Compartilhar" + "share": "Compartilhar", + "refresh": "Atualizar Metadados" }, "lists": { "all": "Todos", @@ -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": { @@ -200,6 +223,9 @@ "targetFormat": "Formato", "defaultBitRate": "Bitrate padrão", "command": "Comando" + }, + "choices": { + "noDefaultBitRate": "Nenhum" } }, "playlist": { @@ -214,7 +240,8 @@ "songCount": "Músicas", "comment": "Comentário", "sync": "Auto-importar", - "path": "Importar de" + "path": "Importar de", + "starred": "Favorita" }, "actions": { "selectPlaylist": "Selecione a playlist:", @@ -301,11 +328,22 @@ "totalDuration": "Duração", "defaultNewUsers": "Padrão para Novos Usuários", "createdAt": "Data de Criação", - "updatedAt": "Últ. Atualização" + "updatedAt": "Últ. Atualização", + "pidAlbum": "Agrupamento de álbuns", + "pidTrack": "Identificação das faixas" }, "sections": { "basic": "Informações Básicas", - "statistics": "Estatísticas" + "statistics": "Estatísticas", + "pid": "IDs Persistentes" + }, + "pid": { + "global": "Usar configuração global (%{value})", + "folder": "Pasta (um álbum por pasta)", + "custom": "Personalizado", + "spec": "Especificação do PID", + "help": "Tags e atributos que identificam um item. Consulte a sintaxe na documentação:", + "docs": "IDs Persistentes" }, "actions": { "scan": "Scanear Biblioteca", @@ -335,7 +373,9 @@ "messages": { "deleteConfirm": "Tem certeza que deseja excluir esta biblioteca? Isso removerá todos os dados associados.", "scanInProgress": "Scan em progresso...", - "noLibrariesAssigned": "Nenhuma biblioteca atribuída a este usuário" + "noLibrariesAssigned": "Nenhuma biblioteca atribuída a este usuário", + "pidChangeTitle": "Alterar os IDs persistentes?", + "pidChangeConfirm": "Ao salvar, os álbuns desta biblioteca serão reagrupados e as faixas serão identificadas novamente. Um scan completo da biblioteca começará imediatamente. As marcações como favoritas, as classificações e as contagens de reprodução das faixas serão mantidas. Os favoritos e as classificações dos álbuns serão transferidos para os novos álbuns quando um álbum antigo corresponder a um novo." } }, "plugin": { @@ -394,7 +434,6 @@ "invalidJson": "A configuração deve ser um JSON válido" }, "messages": { - "idHelp": "O ID do plugin, derivado do nome do arquivo. Use-o ao referenciar este plugin em opções de configuração, como Agents.", "configHelp": "Configure o plugin usando pares chave-valor. Deixe vazio se o plugin não precisa de configuração.", "clickPermissions": "Clique em uma permissão para ver detalhes", "noConfig": "Nenhuma configuração definida", @@ -408,7 +447,8 @@ "requiredHosts": "Hosts necessários", "configValidationError": "Falha na validação da configuração:", "schemaRenderError": "Não foi possível renderizar o formulário de configuração. O schema do plugin pode estar inválido.", - "allowWriteAccessHelp": "Quando habilitado, o plugin pode modificar arquivos nos diretórios das bibliotecas. Por padrão, plugins têm acesso somente leitura." + "allowWriteAccessHelp": "Quando habilitado, o plugin pode modificar arquivos nos diretórios das bibliotecas. Por padrão, plugins têm acesso somente leitura.", + "idHelp": "O ID do plugin, derivado do nome do arquivo. Use-o ao referenciar este plugin em opções de configuração, como Agents." }, "placeholders": { "configKey": "chave", @@ -605,7 +645,10 @@ "coverRemoved": "Capa removida", "coverUploadError": "Erro ao enviar capa", "coverRemoveError": "Erro ao remover capa", - "metadataRefreshStarted": "Atualizando metadados em segundo plano" + "metadataRefreshStarted": "Atualizando metadados em segundo plano", + "quickConnectApproved": "%{app} em %{device} está conectado agora", + "quickConnectInvalidCode": "Código inválido ou expirado", + "quickConnectError": "Não foi possível aprovar o código" }, "menu": { "library": "Biblioteca", @@ -640,6 +683,15 @@ "multipleLibraries": "%{selected} de %{total} Bibliotecas", "selectLibraries": "Selecionar Bibliotecas", "none": "Nenhuma" + }, + "onlyFavourites": "Somente favoritas", + "quickConnect": { + "name": "Conexão Rápida", + "code": "Código", + "help": "Digite o código exibido por um aplicativo Jellyfin para conectá-lo à sua conta", + "confirm": "Conectar %{app} %{version} em %{device} à sua conta?", + "continue": "Continuar", + "approve": "Aprovar" } }, "player": { diff --git a/resources/i18n/sv.json b/resources/i18n/sv.json index 228cc2cf3..ff4f53a6b 100644 --- a/resources/i18n/sv.json +++ b/resources/i18n/sv.json @@ -11,22 +11,22 @@ "title": "Titel", "artist": "Artist", "album": "Album", - "path": "Sökväg", + "path": "Filsökväg", "genre": "Genre", - "compilation": "Samling", + "compilation": "Samlingsalbum", "year": "År", "size": "Filstorlek", "updatedAt": "Uppdaterad", - "bitRate": "Bitrate", - "discSubtitle": "Underrubrik", + "bitRate": "Bithastighet", + "discSubtitle": "Skivans undertitel", "starred": "Favorit", "comment": "Kommentar", "rating": "Betyg", "quality": "Kvalitet", "bpm": "BPM", "playDate": "Senast spelad", - "channels": "Channels", - "createdAt": "Skapad", + "channels": "Kanaler", + "createdAt": "Tillagd", "grouping": "Gruppering", "mood": "Stämning", "participants": "Ytterligare medverkande", @@ -35,19 +35,21 @@ "rawTags": "Omodifierade taggar", "bitDepth": "Bitdjup", "sampleRate": "Samplingsfrekvens", - "missing": "Saknade", + "missing": "Saknas", "libraryName": "Bibliotek", "composer": "Kompositör", - "disc": "Disc %{discNumber}" + "disc": "Skiva %{discNumber}", + "albumGain": "Volymjustering (album)", + "trackGain": "Volymjustering (låt)" }, "actions": { "addToQueue": "Lägg till i kön", "playNow": "Spela nu", "addToPlaylist": "Lägg till i spellista", - "shuffleAll": "Shuffle", + "shuffleAll": "Blanda alla", "download": "Ladda ner", - "playNext": "Spela nästa", - "info": "Mer information", + "playNext": "Spela härnäst", + "info": "Visa information", "showInPlaylist": "Visa i spellista", "instantMix": "Direktmix" } @@ -62,12 +64,12 @@ "playCount": "Spelningar", "name": "Namn", "genre": "Genre", - "compilation": "Samling", + "compilation": "Samlingsalbum", "year": "År", "updatedAt": "Uppdaterad", "comment": "Kommentar", "rating": "Betyg", - "createdAt": "Skapad", + "createdAt": "Tillagt", "size": "Storlek", "originalDate": "Originaldatum", "releaseDate": "Utgivningsdatum", @@ -80,27 +82,28 @@ "media": "Media", "mood": "Stämning", "date": "Inspelningsdatum", - "missing": "Saknade", + "missing": "Saknas", "libraryName": "Bibliotek" }, "actions": { "playAll": "Spela", "playNext": "Spela härnäst", "addToQueue": "Lägg till i kön", - "shuffle": "Shuffle", + "shuffle": "Blanda", "addToPlaylist": "Lägg till i spellista", "download": "Ladda ner", - "info": "Mer information", - "share": "Dela" + "info": "Visa information", + "share": "Dela", + "refresh": "Uppdatera metadata" }, "lists": { "all": "Alla", - "random": "Blanda", - "recentlyAdded": "Senast tillagda", - "recentlyPlayed": "Senast spelade", + "random": "Slumpmässiga", + "recentlyAdded": "Nyligen tillagda", + "recentlyPlayed": "Nyligen spelade", "mostPlayed": "Mest spelade", "starred": "Favoriter", - "topRated": "Bästa betyg" + "topRated": "Högst betygsatta" } }, "artist": { @@ -114,26 +117,26 @@ "genre": "Genre", "size": "Storlek", "role": "Roll", - "missing": "Saknade" + "missing": "Saknas" }, "roles": { "albumartist": "Albumartist |||| Albumartister", "artist": "Artist |||| Artister", - "composer": "Kompositör |||| Kompositörer", + "composer": "Kompositör |||| Kompositörer", "conductor": "Dirigent |||| Dirigenter", "lyricist": "Textförfattare |||| Textförfattare", "arranger": "Arrangör |||| Arrangörer", "producer": "Producent |||| Producenter", "director": "Inspelningsledare |||| Inspelningsledare", "engineer": "Ljudtekniker |||| Ljudtekniker", - "mixer": "Mixare |||| Mixare", - "remixer": "Remixare |||| Remixare", - "djmixer": "DJ-mixare |||| DJ-mixare", - "performer": "Utövande artist |||| Utövande artister", - "maincredit": "Albumartister eller Artist |||| Albumartister eller Artister" + "mixer": "Mixare |||| Mixare", + "remixer": "Remixare |||| Remixare", + "djmixer": "DJ-mixare |||| DJ-mixare", + "performer": "Utövande artist |||| Utövande artister", + "maincredit": "Albumartist eller artist |||| Albumartister eller artister" }, "actions": { - "shuffle": "Shuffle", + "shuffle": "Blanda", "radio": "Radio", "topSongs": "Topplåtar" } @@ -142,7 +145,7 @@ "name": "Användare |||| Användare", "fields": { "userName": "Användarnamn", - "isAdmin": "Är admin", + "isAdmin": "Administratör", "lastLoginAt": "Senaste inloggning", "updatedAt": "Uppdaterad", "name": "Namn", @@ -151,13 +154,15 @@ "changePassword": "Byt lösenord?", "currentPassword": "Nuvarande lösenord", "newPassword": "Nytt lösenord", - "token": "Token", + "token": "Åtkomsttoken", "lastAccessAt": "Senaste åtkomst", - "libraries": "Bibliotek" + "libraries": "Bibliotek", + "scrobbleFilter": "Scrobblingsfilter" }, "helperTexts": { "name": "Ändringar av ditt namn syns först vid nästa inloggning", - "libraries": "Välj ett bibliotek för denna användare eller lämna blankt för standardbibliotek" + "libraries": "Välj vilka bibliotek användaren ska ha tillgång till, eller lämna tomt för att använda standardbiblioteken", + "scrobbleFilter": "Låtar som matchar dessa regler för smarta spellistor skickas inte till Last.fm, ListenBrainz eller scrobblingtillägg. Samma JSON-syntax och beteende som för smarta spellistor används. Exempel: {\"all\":[{\"lt\":{\"rating\":4}}]}. Lämna tomt för att scrobbla allt. Antalet spelningar som lagras lokalt påverkas inte." }, "notifications": { "created": "Användare skapad", @@ -165,13 +170,14 @@ "deleted": "Användare borttagen" }, "message": { - "listenBrainzToken": "Ange din ListenBrainz användar-token.", + "listenBrainzToken": "Ange din användartoken för ListenBrainz.", "clickHereForToken": "Klicka här för att hämta din token", "selectAllLibraries": "Välj alla bibliotek", - "adminAutoLibraries": "Administratörer har automatiskt tillgång till alla bibliotek" + "adminAutoLibraries": "Administratörer har automatiskt åtkomst till alla bibliotek" }, "validation": { - "librariesRequired": "Minst ett bibliotek måste väljas för icke-administratörer" + "librariesRequired": "Minst ett bibliotek måste väljas för användare som inte är administratörer", + "invalidScrobbleFilter": "Måste innehålla giltiga regler för smarta spellistor. Begränsning av antalet låtar, förskjutning och uppdateringsfördröjning stöds inte." } }, "player": { @@ -179,12 +185,12 @@ "fields": { "name": "Namn", "transcodingId": "Omkodning", - "maxBitRate": "Max. bitrate", + "maxBitRate": "Högsta bithastighet", "client": "Klient", "userName": "Användarnamn", "lastSeen": "Senast sedd", - "reportRealPath": "Visa hela sökvägen", - "scrobbleEnabled": "Scrobbla till extern tjänst" + "reportRealPath": "Rapportera den faktiska sökvägen", + "scrobbleEnabled": "Scrobbla till externa tjänster" } }, "transcoding": { @@ -192,8 +198,11 @@ "fields": { "name": "Namn", "targetFormat": "Målformat", - "defaultBitRate": "Standardbitrate", + "defaultBitRate": "Standardbithastighet", "command": "Kommando" + }, + "choices": { + "noDefaultBitRate": "Ingen" } }, "playlist": { @@ -207,8 +216,9 @@ "createdAt": "Skapad", "songCount": "Låtar", "comment": "Kommentar", - "sync": "Auto-import", - "path": "Importera från" + "sync": "Automatisk import", + "path": "Importera från", + "starred": "Favorit" }, "actions": { "selectPlaylist": "Välj en spellista:", @@ -216,24 +226,24 @@ "export": "Exportera", "makePublic": "Gör offentlig", "makePrivate": "Gör privat", - "saveQueue": "Spara kö till spellista", - "searchOrCreate": "Sök spellista eller skapa ny...", - "pressEnterToCreate": "Tryck Enter för att skapa ny spellista", - "removeFromSelection": "Ta bort från urval" + "saveQueue": "Spara kön i en spellista", + "searchOrCreate": "Sök spellistor eller skriv för att skapa en ny...", + "pressEnterToCreate": "Tryck på Enter för att skapa en ny spellista", + "removeFromSelection": "Ta bort från urvalet" }, "message": { - "duplicate_song": "Lägg till dubletter", - "song_exist": "Vissa låtar finns redan i spellistan. Vill du lägga till dubbletterna eller hoppa över dem?", - "noPlaylistsFound": "Hittade inga spellistor", + "duplicate_song": "Lägg till dubbletter", + "song_exist": "Du håller på att lägga till dubbletter i spellistan. Vill du lägga till dem eller hoppa över dem?", + "noPlaylistsFound": "Inga spellistor hittades", "noPlaylists": "Inga spellistor tillgängliga" } }, "radio": { - "name": "Radio |||| Radior", + "name": "Radiostation |||| Radiostationer", "fields": { "name": "Namn", - "streamUrl": "Stream-URL", - "homePageUrl": "Hemside-URL", + "streamUrl": "Strömmens URL", + "homePageUrl": "Webbplatsens URL", "updatedAt": "Uppdaterad", "createdAt": "Skapad" }, @@ -242,7 +252,7 @@ } }, "share": { - "name": "Dela |||| Delningar", + "name": "Delning |||| Delningar", "fields": { "username": "Delad av", "url": "URL", @@ -252,11 +262,13 @@ "lastVisitedAt": "Senast besökt", "visitCount": "Besök", "format": "Format", - "maxBitRate": "Max. bitrate", + "maxBitRate": "Högsta bithastighet", "updatedAt": "Uppdaterad", "createdAt": "Skapad", - "downloadable": "Tillåt nedladdning?" - } + "downloadable": "Tillåt nedladdningar?" + }, + "notifications": {}, + "actions": {} }, "missing": { "name": "Saknad fil |||| Saknade filer", @@ -267,11 +279,11 @@ "libraryName": "Bibliotek" }, "actions": { - "remove": "Radera", - "remove_all": "Radera alla" + "remove": "Ta bort", + "remove_all": "Ta bort alla" }, "notifications": { - "removed": "Saknade fil(er) borttagna" + "removed": "Saknade filer har tagits bort" }, "empty": "Inga saknade filer" }, @@ -280,8 +292,8 @@ "fields": { "name": "Namn", "path": "Sökväg", - "remotePath": "Ta bort sökväg", - "lastScanAt": "Senaste scan", + "remotePath": "Fjärrsökväg", + "lastScanAt": "Senaste skanning", "songCount": "Låtar", "albumCount": "Album", "artistCount": "Artister", @@ -302,33 +314,33 @@ "statistics": "Statistik" }, "actions": { - "scan": "Scanna bibliotek", + "scan": "Skanna biblioteket", "manageUsers": "Hantera användaråtkomst", - "viewDetails": "Se detaljer", - "quickScan": "Snabbscan", - "fullScan": "Komplett scan" + "viewDetails": "Visa detaljer", + "quickScan": "Snabbskanning", + "fullScan": "Fullständig skanning" }, "notifications": { "created": "Biblioteket har skapats", "updated": "Biblioteket har uppdaterats", - "deleted": "Biblioteket har raderats", - "scanStarted": "Biblioteksscan startad", - "scanCompleted": "Biblioteksscan avslutad", - "quickScanStarted": "Snabbscan startad", - "fullScanStarted": "Komplett scan startad", - "scanError": "Fel vid start av scan. Se loggarna" + "deleted": "Biblioteket har tagits bort", + "scanStarted": "Skanning av biblioteket har startat", + "scanCompleted": "Skanning av biblioteket är klar", + "quickScanStarted": "Snabbskanning har startat", + "fullScanStarted": "Fullständig skanning har startat", + "scanError": "Kunde inte starta skanningen. Kontrollera loggarna" }, "validation": { - "nameRequired": "Biblioteksnamn krävs", - "pathRequired": "Bibliotekssökväg krävs", - "pathNotDirectory": "Bibliotekssökvägen måste vara en katalog", - "pathNotFound": "Bibliotekssökväg hittades inte", - "pathNotAccessible": "Bibliotekssökväg inte tillgänglig", - "pathInvalid": "Ogiltig bibliotekssökväg" + "nameRequired": "Ange ett biblioteksnamn", + "pathRequired": "Ange en sökväg till biblioteket", + "pathNotDirectory": "Bibliotekets sökväg måste peka på en katalog", + "pathNotFound": "Bibliotekets sökväg hittades inte", + "pathNotAccessible": "Bibliotekets sökväg är inte åtkomlig", + "pathInvalid": "Ogiltig sökväg till biblioteket" }, "messages": { - "deleteConfirm": "Är du säker på att du vill ta bort detta bibliotek? Detta raderar all förbunden data och användartillgång.", - "scanInProgress": "Scanning pågår...", + "deleteConfirm": "Är du säker på att du vill ta bort det här biblioteket? Alla tillhörande data och användarnas åtkomst till biblioteket tas bort.", + "scanInProgress": "Skanning pågår...", "noLibrariesAssigned": "Inga bibliotek har tilldelats den här användaren" } }, @@ -339,44 +351,44 @@ "name": "Namn", "description": "Beskrivning", "version": "Version", - "author": "Författare", - "website": "Website", + "author": "Upphovsperson", + "website": "Webbplats", "permissions": "Behörigheter", - "enabled": "Aktiverad", + "enabled": "Aktiverat", "status": "Status", "path": "Sökväg", "lastError": "Fel", "hasError": "Fel", - "updatedAt": "Uppdaterad", - "createdAt": "Installerad", + "updatedAt": "Uppdaterat", + "createdAt": "Installerat", "configKey": "Nyckel", "configValue": "Värde", "allUsers": "Tillåt alla användare", "selectedUsers": "Valda användare", "allLibraries": "Tillåt alla bibliotek", "selectedLibraries": "Valda bibliotek", - "allowWriteAccess": "Tillåt skrivrättigheter" + "allowWriteAccess": "Tillåt skrivåtkomst" }, "sections": { "status": "Status", - "info": "Tilläggsinformation", + "info": "Information om tillägget", "configuration": "Konfiguration", "manifest": "Manifest", - "usersPermission": "Användarbehörigheter", - "libraryPermission": "Biblioteksbehörigheter" + "usersPermission": "Åtkomst till användare", + "libraryPermission": "Åtkomst till bibliotek" }, "status": { - "enabled": "Aktiverad", - "disabled": "Inaktiverad" + "enabled": "Aktiverat", + "disabled": "Inaktiverat" }, "actions": { "enable": "Aktivera", "disable": "Inaktivera", - "disabledDueToError": "Åtgärda felet innan aktivering", + "disabledDueToError": "Åtgärda felet före aktivering", "disabledUsersRequired": "Välj användare före aktivering", "disabledLibrariesRequired": "Välj bibliotek före aktivering", "addConfig": "Lägg till konfiguration", - "rescan": "Scanna om" + "rescan": "Skanna om" }, "notifications": { "enabled": "Tillägg aktiverat", @@ -391,17 +403,18 @@ "configHelp": "Konfigurera tillägget med nyckel–värde-par. Lämna tomt om tillägget inte kräver någon konfiguration.", "clickPermissions": "Klicka på en behörighet för mer information", "noConfig": "Ingen konfiguration angiven", - "allUsersHelp": "När den är aktiverad får tillägget tillgång till alla användare, inklusive de som skapas i framtiden.", + "allUsersHelp": "När alternativet är aktiverat får tillägget åtkomst till alla användare, även de som skapas i framtiden.", "noUsers": "Inga användare valda", "permissionReason": "Orsak", - "usersRequired": "Detta tillägg kräver åtkomst till användarinformation. Välj vilka användare insticksprogrammet ska ha åtkomst till, eller aktivera 'Tillåt alla användare'.", - "allLibrariesHelp": "När den är aktiverad får tillägget tillgång till alla bibliotek, inklusive de som skapas i framtiden.", + "usersRequired": "Det här tillägget behöver åtkomst till användarinformation. Välj vilka användare tillägget får åtkomst till eller aktivera 'Tillåt alla användare'.", + "allLibrariesHelp": "När alternativet är aktiverat får tillägget åtkomst till alla bibliotek, även de som skapas i framtiden.", "noLibraries": "Inga bibliotek valda", - "librariesRequired": "Detta tillägg kräver tillgång till biblioteksinformation. Välj vilka bibliotek tillägget kan komma åt eller aktivera 'Tillåt alla bibliotek'.", - "requiredHosts": "Krävda värdar", + "librariesRequired": "Det här tillägget behöver åtkomst till biblioteksinformation. Välj vilka bibliotek tillägget får åtkomst till eller aktivera 'Tillåt alla bibliotek'.", + "requiredHosts": "Värdar som krävs", "configValidationError": "Validering av konfigurationen misslyckades:", - "schemaRenderError": "Kunde inte rendera konfigurationsformuläret. Tilläggets schema kan vara ogiltigt.", - "allowWriteAccessHelp": "När detta är aktiverat kan tillägget ändra filer i bibliotekets kataloger. Som standard har tillägget endast läsrättigheter." + "schemaRenderError": "Kunde inte visa konfigurationsformuläret. Tilläggets schema kan vara ogiltigt.", + "allowWriteAccessHelp": "När alternativet är aktiverat kan tillägget ändra filer i bibliotekens kataloger. Som standard har tillägg endast läsåtkomst.", + "idHelp": "Tilläggets ID hämtas från filnamnet. Använd det när du hänvisar till tillägget i konfigurationsalternativ, till exempel Agents." }, "placeholders": { "configKey": "nyckel", @@ -412,40 +425,40 @@ "ra": { "auth": { "welcome1": "Tack för att du installerade Navidrome!", - "welcome2": "Skapa först ett admin-konto", + "welcome2": "Börja med att skapa ett administratörskonto", "confirmPassword": "Bekräfta lösenord", - "buttonCreateAdmin": "Skapa admin-konto", + "buttonCreateAdmin": "Skapa administratörskonto", "auth_check_error": "Logga in för att fortsätta", "user_menu": "Profil", "username": "Användarnamn", "password": "Lösenord", "sign_in": "Logga in", - "sign_in_error": "Felaktig inloggning, försök igen", + "sign_in_error": "Inloggningen misslyckades. Försök igen", "logout": "Logga ut", - "insightsCollectionNote": "Navidrome samlar anonym användardata för att\nhjälpa projektet att bli bättre. Klicka [här]\nför att läsa mer och avaktivera om du vill" + "insightsCollectionNote": "Navidrome samlar in anonyma användningsdata för\natt förbättra projektet. Klicka [här] för att läsa mer\noch välja bort insamlingen" }, "validation": { "invalidChars": "Använd enbart bokstäver och siffror", - "passwordDoesNotMatch": "Lösenordet matchar inte", - "required": "Krävs", - "minLength": "Måste ha minst %{min} tecken", - "maxLength": "Får maximalt ha %{max} tecken", + "passwordDoesNotMatch": "Lösenorden stämmer inte överens", + "required": "Obligatoriskt", + "minLength": "Måste innehålla minst %{min} tecken", + "maxLength": "Får innehålla högst %{max} tecken", "minValue": "Måste vara minst %{min}", - "maxValue": "Får maximalt vara %{max}", - "number": "Måste vara ett nummer", + "maxValue": "Får vara högst %{max}", + "number": "Måste vara ett tal", "email": "Måste vara en giltig e-postadress", - "oneOf": "Måste vara en av: %{options}", - "regex": "Måste matcha ett specifikt format (regexp): %{pattern}", - "unique": "Måste vara unik", + "oneOf": "Måste vara något av: %{options}", + "regex": "Måste matcha det reguljära uttrycket: %{pattern}", + "unique": "Måste vara unikt", "url": "Måste vara en giltig URL" }, "action": { "add_filter": "Lägg till filter", "add": "Lägg till", "back": "Tillbaka", - "bulk_actions": "1 objekt vald |||| %{smart_count} objekt valda", + "bulk_actions": "1 objekt valt |||| %{smart_count} objekt valda", "cancel": "Avbryt", - "clear_input_value": "Rensa", + "clear_input_value": "Rensa värdet", "clone": "Klona", "confirm": "Bekräfta", "create": "Skapa", @@ -454,8 +467,8 @@ "export": "Exportera", "list": "Lista", "refresh": "Uppdatera", - "remove_filter": "Ta bort filter", - "remove": "Radera", + "remove_filter": "Ta bort det här filtret", + "remove": "Ta bort", "save": "Spara", "search": "Sök", "show": "Visa", @@ -477,15 +490,15 @@ }, "page": { "create": "Skapa %{name}", - "dashboard": "Dashboard", + "dashboard": "Översikt", "edit": "%{name} #%{id}", "error": "Ett fel uppstod", "list": "%{name}", - "loading": "Laddar", - "not_found": "Hittade inget", + "loading": "Läser in", + "not_found": "Hittades inte", "show": "%{name} #%{id}", - "empty": "Ingen %{name} ännu.", - "invite": "Vill du lägga till en?" + "empty": "%{name}: inga poster ännu.", + "invite": "Vill du lägga till något?" }, "input": { "file": { @@ -497,13 +510,13 @@ "upload_single": "Dra och släpp en bild som ska laddas upp eller klicka för att välja en bild." }, "references": { - "all_missing": "Hittade ingen referensdata.", - "many_missing": "Minst en av de associerade referenserna verkar inte längre vara tillgänglig.", - "single_missing": "Associerade referenser verkar inte längre vara tillgängliga." + "all_missing": "Kunde inte hitta referensdata.", + "many_missing": "Minst en av de kopplade referenserna verkar inte längre vara tillgänglig.", + "single_missing": "Den kopplade referensen verkar inte längre vara tillgänglig." }, "password": { - "toggle_visible": "Dölj password", - "toggle_hidden": "Visa password" + "toggle_visible": "Dölj lösenord", + "toggle_hidden": "Visa lösenord" } }, "message": { @@ -511,16 +524,16 @@ "are_you_sure": "Är du säker?", "bulk_delete_content": "Vill du verkligen ta bort %{name}? |||| Vill du verkligen ta bort dessa %{smart_count} objekt?", "bulk_delete_title": "Ta bort %{name} |||| Ta bort %{smart_count} %{name}", - "delete_content": "Vill du verkligen ta bort detta innehåll?", + "delete_content": "Vill du verkligen ta bort det här objektet?", "delete_title": "Ta bort %{name} #%{id}", "details": "Detaljer", "error": "Ett klientfel uppstod och begäran kunde inte slutföras.", - "invalid_form": "Formuläret är ogiltigt. Kontrollera eventuella fel", - "loading": "Sidan läses in, var god vänta", + "invalid_form": "Formuläret innehåller fel. Kontrollera uppgifterna", + "loading": "Sidan läses in. Vänta ett ögonblick", "no": "Nej", "not_found": "Antingen skrev du fel URL eller så följde du en ogiltig länk.", "yes": "Ja", - "unsaved_changes": "Du har osparade ändringar. Ignorera dem?" + "unsaved_changes": "Vissa ändringar har inte sparats. Vill du verkligen ignorera dem?" }, "navigation": { "no_results": "Inga resultat hittades", @@ -529,23 +542,23 @@ "page_out_from_end": "Det finns inga fler sidor", "page_out_from_begin": "Det finns ingen sida före sida 1", "page_range_info": "%{offsetBegin}-%{offsetEnd} av %{total}", - "page_rows_per_page": "Antal per sida:", + "page_rows_per_page": "Objekt per sida:", "next": "Nästa", "prev": "Föregående", "skip_nav": "Hoppa till innehåll" }, "notification": { - "updated": "Element uppdaterat |||| %{smart_count} element uppdaterade", - "created": "Element skapat", - "deleted": "Element borttaget |||| %{smart_count} element borttagna", - "bad_item": "Felaktigt element", - "item_doesnt_exist": "Element finns inte", + "updated": "Objekt uppdaterat |||| %{smart_count} objekt uppdaterade", + "created": "Objekt skapat", + "deleted": "Objekt borttaget |||| %{smart_count} objekt borttagna", + "bad_item": "Felaktigt objekt", + "item_doesnt_exist": "Objektet finns inte", "http_error": "Kommunikationsfel med servern", - "data_provider_error": "Fel i dataProvider. Kontrollera din konsol för mer information.", - "i18n_error": "Kunde inte läsa in översättningen av det valda språket", + "data_provider_error": "Fel i dataProvider. Kontrollera webbläsarens konsol för mer information.", + "i18n_error": "Kunde inte läsa in översättningarna för det valda språket", "canceled": "Åtgärden avbröts", - "logged_out": "Sessionen har avslutats, anslut på nytt.", - "new_version": "Det finns en ny version! Uppdatera detta fönster." + "logged_out": "Sessionen har avslutats. Logga in igen.", + "new_version": "En ny version finns tillgänglig! Ladda om det här fönstret." }, "toggleFieldsMenu": { "columnsToDisplay": "Kolumner att visa", @@ -556,39 +569,39 @@ }, "message": { "note": "OBSERVERA", - "transcodingDisabled": "Inställning för kodning via webbgränssnittet är av säkerhetsskäl ej aktiverat. Starta om servern med alternativet %{config} markerat om du vill göra ändringar (redigera eller lägga till).", - "transcodingEnabled": "Navidrome körs för närvarande med %{config}, vilket gör att systemkommandon kan köras från webbplattformen. Du rekommenderas av säkerhetsskäl att du stänger av den och bara slår på den när du ställer in omkodning.", - "songsAddedToPlaylist": "La till en låt i spellistan |||| La till %{smart_count} låtar i spellistan", - "noPlaylistsAvailable": "Ingen tillgänglig", - "delete_user_title": "Ta bort användare '%{name}'", - "delete_user_content": "Är du säker på att du vill ta bort denna användare (inklusive spellistor och inställningar)?", - "notifications_blocked": "Du har blockerat meddelanden från denna sajt in din webbläsares inställningar", - "notifications_not_available": "Denna webbläsare stödjer inte skrivbordsmeddelanden eller du använder inte Navidrome via https", - "lastfmLinkSuccess": "Last.fm är länkat och scrobbling är aktivt", - "lastfmLinkFailure": "Last.fm kunde inte länkas", - "lastfmUnlinkSuccess": "Last.fm är inte längre länkat och scrobbling är deaktiverat", - "lastfmUnlinkFailure": "Last.fm kunde inte avlänkas", + "transcodingDisabled": "Av säkerhetsskäl går det inte att ändra omkodningsinställningarna i webbgränssnittet. Starta om servern med konfigurationsalternativet %{config} för att redigera eller lägga till omkodningsinställningar.", + "transcodingEnabled": "Navidrome körs med %{config}. Det gör det möjligt att köra systemkommandon via omkodningsinställningarna i webbgränssnittet. Av säkerhetsskäl rekommenderar vi att alternativet är inaktiverat och bara aktiveras när omkodningen ska konfigureras.", + "songsAddedToPlaylist": "1 låt har lagts till i spellistan |||| %{smart_count} låtar har lagts till i spellistan", + "noPlaylistsAvailable": "Inga tillgängliga", + "delete_user_title": "Ta bort användaren '%{name}'", + "delete_user_content": "Är du säker på att du vill ta bort den här användaren och alla användarens data, inklusive spellistor och inställningar?", + "notifications_blocked": "Du har blockerat aviseringar från den här webbplatsen i webbläsarens inställningar", + "notifications_not_available": "Webbläsaren stöder inte skrivbordsaviseringar, eller så ansluter du inte till Navidrome via HTTPS", + "lastfmLinkSuccess": "Last.fm har kopplats och scrobbling har aktiverats", + "lastfmLinkFailure": "Det gick inte att koppla Last.fm", + "lastfmUnlinkSuccess": "Kopplingen till Last.fm har tagits bort och scrobbling har inaktiverats", + "lastfmUnlinkFailure": "Det gick inte att ta bort kopplingen till Last.fm", "openIn": { "lastfm": "Öppna i Last.fm", "musicbrainz": "Öppna i MusicBrainz" }, "lastfmLink": "Läs mer...", - "listenBrainzLinkSuccess": "ListenBrainz är länkat och scrobbling är aktivt som användare: %{user}", - "listenBrainzLinkFailure": "ListenBrainz kunde inte länkas: %{error}", - "listenBrainzUnlinkSuccess": "ListenBrainz är inte längre länkat och scrobbling är deaktiverat", - "listenBrainzUnlinkFailure": "ListenBrainz kunde inte avlänkas", + "listenBrainzLinkSuccess": "ListenBrainz har kopplats och scrobbling har aktiverats för användaren: %{user}", + "listenBrainzLinkFailure": "Det gick inte att koppla ListenBrainz: %{error}", + "listenBrainzUnlinkSuccess": "Kopplingen till ListenBrainz har tagits bort och scrobbling har inaktiverats", + "listenBrainzUnlinkFailure": "Det gick inte att ta bort kopplingen till ListenBrainz", "downloadOriginalFormat": "Ladda ner i originalformat", "shareOriginalFormat": "Dela i originalformat", "shareDialogTitle": "Dela %{resource} '%{name}'", - "shareBatchDialogTitle": "Dela en %{resource} |||| Dela %{smart_count} %{resource}", - "shareSuccess": "URL kopierades till urklipp: %{url}", - "shareFailure": "Fel vid kopiering av URL %{url} till urklipp", + "shareBatchDialogTitle": "Dela 1 %{resource} |||| Dela %{smart_count} %{resource}", + "shareSuccess": "URL:en har kopierats till urklipp: %{url}", + "shareFailure": "Kunde inte kopiera URL:en %{url} till urklipp", "downloadDialogTitle": "Ladda ner %{resource} '%{name}' (%{size})", "shareCopyToClipboard": "Kopiera till urklipp: Ctrl+C, Enter", "remove_missing_title": "Ta bort saknade filer", - "remove_missing_content": "Är du säker på att du vill ta bort de valda saknade filerna från databasen? Detta kommer permanent radera alla referenser till dem, inklusive antal spelningar och betyg.", + "remove_missing_content": "Är du säker på att du vill ta bort de valda saknade filerna från databasen? Alla referenser till dem, inklusive antal spelningar och betyg, tas bort permanent.", "remove_all_missing_title": "Ta bort alla saknade filer", - "remove_all_missing_content": "Är du säker på att du vill ta bort alla saknade filer från databasen? Detta kommer permanent radera alla referenser till dem, inklusive antal spelningar och betyg.", + "remove_all_missing_content": "Är du säker på att du vill ta bort alla saknade filer från databasen? Alla referenser till dem, inklusive antal spelningar och betyg, tas bort permanent.", "noSimilarSongsFound": "Hittade inga liknande låtar", "noTopSongsFound": "Hittade inga topplåtar", "startingInstantMix": "Laddar direktmix...", @@ -597,7 +610,11 @@ "coverUploaded": "Omslagsbild uppdaterad", "coverRemoved": "Omslagsbild borttagen", "coverUploadError": "Fel vid uppladdning av omslagsbild", - "coverRemoveError": "Fel vid borttagning av omslagsbild" + "coverRemoveError": "Fel vid borttagning av omslagsbild", + "metadataRefreshStarted": "Metadata uppdateras i bakgrunden", + "quickConnectApproved": "%{app} på %{device} är nu inloggad", + "quickConnectInvalidCode": "Koden är ogiltig eller har gått ut", + "quickConnectError": "Kunde inte godkänna koden" }, "menu": { "library": "Bibliotek", @@ -610,17 +627,17 @@ "theme": "Tema", "language": "Språk", "defaultView": "Standardvy", - "desktop_notifications": "Skrivbordsmeddelanden", + "desktop_notifications": "Skrivbordsaviseringar", "lastfmScrobbling": "Scrobbla till Last.fm", "listenBrainzScrobbling": "Scrobbla till ListenBrainz", "replaygain": "ReplayGain-läge", - "preAmp": "ReplayGain PreAmp (dB)", + "preAmp": "Förförstärkning för ReplayGain (dB)", "gain": { "none": "Inaktiverad", - "album": "Använd gain för album", - "track": "Använd gain für låtar" + "album": "Använd albumets volymjustering", + "track": "Använd låtens volymjustering" }, - "lastfmNotConfigured": "Last.fm API-nyckel är inte konfigurerad" + "lastfmNotConfigured": "API-nyckeln för Last.fm är inte konfigurerad" } }, "albumList": "Album", @@ -632,10 +649,19 @@ "multipleLibraries": "%{selected} av %{total} bibliotek", "selectLibraries": "Välj bibliotek", "none": "Inga" - } + }, + "quickConnect": { + "name": "Snabbanslutning", + "code": "Kod", + "help": "Ange koden som visas i en Jellyfin-app för att logga in på ditt konto i appen", + "confirm": "Vill du logga in med ditt konto i %{app} %{version} på %{device}?", + "continue": "Fortsätt", + "approve": "Godkänn" + }, + "onlyFavourites": "Visa endast favoriter" }, "player": { - "playListsText": "Spela kön", + "playListsText": "Uppspelningskö", "openText": "Öppna", "closeText": "Stäng", "notContentText": "Ingen musik", @@ -645,26 +671,26 @@ "previousTrackText": "Föregående låt", "reloadText": "Ladda om", "volumeText": "Volym", - "toggleLyricText": "Låttext av/på", + "toggleLyricText": "Visa eller dölj låttexten", "toggleMiniModeText": "Minimera", - "destroyText": "Radera", + "destroyText": "Stäng spelaren och rensa kön", "downloadText": "Ladda ner", - "removeAudioListsText": "Ta bort audiolistor", + "removeAudioListsText": "Rensa kön", "clickToDeleteText": "Klicka för att ta bort %{name}", "emptyLyricText": "Ingen låttext", "playModeText": { "order": "I ordningsföljd", - "orderLoop": "Upprepa", - "singleLoop": "Upprepa en", - "shufflePlay": "Shuffle" + "orderLoop": "Upprepa alla", + "singleLoop": "Upprepa en låt", + "shufflePlay": "Blanda" } }, "about": { "links": { - "homepage": "Hemsida", + "homepage": "Webbplats", "source": "Källkod", - "featureRequests": "Funktionalitetförfrågan", - "lastInsightsCollection": "Senaste Insights-kollektion", + "featureRequests": "Förslag på nya funktioner", + "lastInsightsCollection": "Senaste insamling av användningsdata", "insights": { "disabled": "Inaktiverad", "waiting": "Väntar" @@ -672,16 +698,16 @@ }, "tabs": { "about": "Om", - "config": "Inställningar" + "config": "Konfiguration" }, "config": { - "configName": "Inställningsnamn", + "configName": "Inställning", "environmentVariable": "Miljövariabel", - "currentValue": "Nuvarande värde", - "configurationFile": "Inställningsfil", - "exportToml": "Exportera inställningar (TOML)", - "exportSuccess": "Inställningarna kopierade till urklippet i TOML-format", - "exportFailed": "Kopiering av inställningarna misslyckades", + "currentValue": "Aktuellt värde", + "configurationFile": "Konfigurationsfil", + "exportToml": "Exportera konfiguration (TOML)", + "exportSuccess": "Konfigurationen har kopierats till urklipp i TOML-format", + "exportFailed": "Kunde inte kopiera konfigurationen", "devFlagsHeader": "Utvecklingsflaggor (kan ändras eller tas bort)", "devFlagsComment": "Dessa inställningar är experimentella och kan tas bort i framtida versioner", "downloadToml": "Ladda ner konfiguration (TOML)" @@ -689,28 +715,28 @@ }, "activity": { "title": "Aktivitet", - "totalScanned": "Genomsökta mappar", - "quickScan": "Snabbscan", - "fullScan": "Komplett scan", - "serverUptime": "Serverdrifttid", + "totalScanned": "Skannade mappar totalt", + "quickScan": "Snabb", + "fullScan": "Fullständig", + "serverUptime": "Serverns drifttid", "serverDown": "OFFLINE", - "scanType": "Typ", - "status": "Fel vid scanning", - "elapsedTime": "Spelad tid", - "selectiveScan": "Urval" + "scanType": "Senaste skanning", + "status": "Skanningsfel", + "elapsedTime": "Förfluten tid", + "selectiveScan": "Selektiv" }, "help": { - "title": "Navidrome kortkommandon", + "title": "Kortkommandon i Navidrome", "hotkeys": { "show_help": "Visa denna hjälp", - "toggle_menu": "Växla sidomeny", + "toggle_menu": "Visa eller dölj sidomenyn", "toggle_play": "Spela / pausa", "prev_song": "Föregående låt", "next_song": "Nästa låt", - "vol_up": "Volym upp", - "vol_down": "Volym ner", - "toggle_love": "Lägg till låt i favoriter", - "current_song": "Hoppa till nuvarande låt" + "vol_up": "Höj volymen", + "vol_down": "Sänk volymen", + "toggle_love": "Lägg till låten i favoriter", + "current_song": "Gå till aktuell låt" } }, "nowPlaying": { @@ -718,4 +744,4 @@ "empty": "Inget spelas", "minutesAgo": "%{smart_count} minut sedan |||| %{smart_count} minuter sedan" } -} \ No newline at end of file +} diff --git a/resources/i18n/th.json b/resources/i18n/th.json index fde89494e..3c5c9e07b 100644 --- a/resources/i18n/th.json +++ b/resources/i18n/th.json @@ -93,7 +93,8 @@ "addToPlaylist": "เพิ่มลงในเพลย์ลิสต์", "download": "ดาวน์โหลด", "info": "ดูรายละเอียด", - "share": "แบ่งปัน" + "share": "แบ่งปัน", + "refresh": "" }, "lists": { "all": "ทั้งหมด", @@ -155,11 +156,13 @@ "newPassword": "รหัสผ่านใหม่", "token": "โทเคน", "lastAccessAt": "เข้าใช้ล่าสุด", - "libraries": "ห้องสมุด" + "libraries": "ห้องสมุด", + "scrobbleFilter": "" }, "helperTexts": { "name": "การเปลี่ยนชื่อจะมีผลในการล็อกอินครั้งถัดไป", - "libraries": "เลือกห้องสมุดสำหรับผู้ใช้นี้หรือปล่อยว่างเพื่อใช้ห้องสมุดเริ่มต้น" + "libraries": "เลือกห้องสมุดสำหรับผู้ใช้นี้หรือปล่อยว่างเพื่อใช้ห้องสมุดเริ่มต้น", + "scrobbleFilter": "" }, "notifications": { "created": "สร้างชื่อผู้ใช้", @@ -173,7 +176,8 @@ "adminAutoLibraries": "ผู้ดูแลเข้าถึงห้องสมุดทั้งหมดโดยอัตโนมัติ" }, "validation": { - "librariesRequired": "ต้องเลือกห้องสมุด 1 ห้อง สำหรับผู้ใช้ที่ไม่ใช่ผู้ดูแล" + "librariesRequired": "ต้องเลือกห้องสมุด 1 ห้อง สำหรับผู้ใช้ที่ไม่ใช่ผู้ดูแล", + "invalidScrobbleFilter": "" } }, "player": { @@ -196,6 +200,9 @@ "targetFormat": "ชนิดไฟล์เสียง", "defaultBitRate": "บิตเรท", "command": "คำสั่ง" + }, + "choices": { + "noDefaultBitRate": "" } }, "playlist": { @@ -210,7 +217,8 @@ "songCount": "เพลง", "comment": "ความคิดเห็น", "sync": "นำเข้าอัตโนมัติ", - "path": "นำเข้าจาก" + "path": "นำเข้าจาก", + "starred": "ชื่นชอบ" }, "actions": { "selectPlaylist": "เลือกเพลย์ลิสต์", @@ -403,7 +411,8 @@ "requiredHosts": "ต้องการ Host", "configValidationError": "การตั้งค่าเกิดความผิดพลาด", "schemaRenderError": "ไม่สามารถแสดงหน้าจอการตั้งค่า อาจเกิดจากความผิดพลาดจากปลั๊กอิน", - "allowWriteAccessHelp": "เมื่อเปิดใช้งาน ปลั๊กอินสามารถแก้ไขไฟล์ในห้องสมุด ปลั๊กอินอยู่ในโหมดอ่านอย่างเดียวเป็นค่าเริ่มต้น" + "allowWriteAccessHelp": "เมื่อเปิดใช้งาน ปลั๊กอินสามารถแก้ไขไฟล์ในห้องสมุด ปลั๊กอินอยู่ในโหมดอ่านอย่างเดียวเป็นค่าเริ่มต้น", + "idHelp": "" }, "placeholders": { "configKey": "คีย์", @@ -599,7 +608,8 @@ "coverUploaded": "ภาพหน้าปกถูกอัพเดทแล้ว", "coverRemoved": "ภาพหน้าปกถูกลบแล้ว", "coverUploadError": "อัพโหลดภาพหน้าปกผิดพลาด", - "coverRemoveError": "ลบภาพหน้าปกผิดพลาด" + "coverRemoveError": "ลบภาพหน้าปกผิดพลาด", + "metadataRefreshStarted": "" }, "menu": { "library": "ห้องสมุดเพลง", @@ -634,7 +644,8 @@ "multipleLibraries": "%{selected} ของ %{total} ห้องสมุด", "selectLibraries": "เลือกห้องสมุด", "none": "ไม่มี" - } + }, + "onlyFavourites": "แสดงเฉพาะที่ชื่นชอบ" }, "player": { "playListsText": "คิวเล่น", diff --git a/resources/i18n/uk.json b/resources/i18n/uk.json index c5644fde7..3049f9840 100644 --- a/resources/i18n/uk.json +++ b/resources/i18n/uk.json @@ -38,7 +38,9 @@ "missing": "Поле відсутнє", "libraryName": "Бібліотека", "composer": "Композитор", - "disc": "Диск %{discNumber}" + "disc": "Диск %{discNumber}", + "albumGain": "", + "trackGain": "" }, "actions": { "addToQueue": "Прослухати пізніше", @@ -91,7 +93,8 @@ "addToPlaylist": "Додати у список відтворення", "download": "Завантажити", "info": "Отримати інформацію", - "share": "Поширити" + "share": "Поширити", + "refresh": "" }, "lists": { "all": "Усі", @@ -153,11 +156,13 @@ "newPassword": "Новий пароль", "token": "Токен", "lastAccessAt": "Останній доступ", - "libraries": "Бібліотеки" + "libraries": "Бібліотеки", + "scrobbleFilter": "" }, "helperTexts": { "name": "Змінене ім'я буде відображатися при наступній авторизації", - "libraries": "Виберіть конкретні бібліотеки для цього користувача, або залиште поле порожнім, щоб використовувати бібліотеки за замовчуванням" + "libraries": "Виберіть конкретні бібліотеки для цього користувача, або залиште поле порожнім, щоб використовувати бібліотеки за замовчуванням", + "scrobbleFilter": "" }, "notifications": { "created": "Користувача створено", @@ -171,7 +176,8 @@ "adminAutoLibraries": "Користувачі-адміністратори автоматично отримують доступ до всіх бібліотек" }, "validation": { - "librariesRequired": "Для користувачів, які не є адміністраторами, має бути обрана хоча б одна бібліотека" + "librariesRequired": "Для користувачів, які не є адміністраторами, має бути обрана хоча б одна бібліотека", + "invalidScrobbleFilter": "" } }, "player": { @@ -194,6 +200,9 @@ "targetFormat": "Цільовий формат", "defaultBitRate": "Швидкість передачі бітів за замовчуванням", "command": "Команда" + }, + "choices": { + "noDefaultBitRate": "" } }, "playlist": { @@ -208,7 +217,8 @@ "songCount": "Пісні", "comment": "Коментар", "sync": "Автоімпорт", - "path": "Імпортувати із" + "path": "Імпортувати із", + "starred": "Улюблене" }, "actions": { "selectPlaylist": "Вибрати список відтворення:", @@ -401,7 +411,8 @@ "requiredHosts": "Обов'язкові хости", "configValidationError": "Перевірка конфігурації зазнала невдачі:", "schemaRenderError": "Неможливо відобразити форму конфігурації. Схема плагіна може бути недійсною.", - "allowWriteAccessHelp": "При включенні плагін може змінювати файли в каталогах бібліотеки. За замовчуванням плагіни мають доступ лише для читання." + "allowWriteAccessHelp": "При включенні плагін може змінювати файли в каталогах бібліотеки. За замовчуванням плагіни мають доступ лише для читання.", + "idHelp": "" }, "placeholders": { "configKey": "ключ", @@ -597,7 +608,8 @@ "coverUploaded": "Обкладинку оновлено", "coverRemoved": "Обкладинка видалена", "coverUploadError": "Помилка завантаження обкладинки", - "coverRemoveError": "Помилка видалення обкладинки" + "coverRemoveError": "Помилка видалення обкладинки", + "metadataRefreshStarted": "" }, "menu": { "library": "Бібліотека", @@ -632,7 +644,8 @@ "multipleLibraries": "%{selected} з %{total} Бібліотеки", "selectLibraries": "Вибір бібліотек", "none": "Відсутня" - } + }, + "onlyFavourites": "Показати улюблене" }, "player": { "playListsText": "Грати по черзі", diff --git a/resources/i18n/zh-Hans.json b/resources/i18n/zh-Hans.json index 24f55357a..cae42460b 100644 --- a/resources/i18n/zh-Hans.json +++ b/resources/i18n/zh-Hans.json @@ -93,7 +93,8 @@ "shuffle": "随机播放", "addToPlaylist": "添加到歌单", "download": "下载", - "info": "查看信息" + "info": "查看信息", + "refresh": "刷新元数据" }, "lists": { "all": "全部", @@ -155,11 +156,13 @@ "currentPassword": "当前密码", "newPassword": "新密码", "token": "令牌", - "libraries": "媒体库" + "libraries": "媒体库", + "scrobbleFilter": "个性化记录过滤器" }, "helperTexts": { "name": "名称的更改将在下次登录时生效", - "libraries": "为此用户选择指定媒体库,留空则使用默认媒体库" + "libraries": "为此用户选择指定媒体库,留空则使用默认媒体库", + "scrobbleFilter": "符合这些智能歌单规则的歌曲不会发送到 Last.fm、ListenBrainz 或个性化记录插件。使用与智能歌单相同的 JSON 语法和行为。例如:{\"all\":[{\"lt\":{\"rating\":4}}]}。留空则记录所有播放。本地播放次数不受影响。" }, "notifications": { "created": "用户已创建", @@ -167,7 +170,8 @@ "deleted": "用户已删除" }, "validation": { - "librariesRequired": "至少为非管理员用户选择一个媒体库" + "librariesRequired": "至少为非管理员用户选择一个媒体库", + "invalidScrobbleFilter": "必须是有效的智能歌单规则。不支持数量限制(limit)、偏移量(offset)和刷新延迟(refresh delay)。" }, "message": { "listenBrainzToken": "输入您的 ListenBrainz 用户令牌", @@ -196,6 +200,9 @@ "targetFormat": "目标格式", "defaultBitRate": "默认比特率", "command": "命令" + }, + "choices": { + "noDefaultBitRate": "无" } }, "playlist": { @@ -393,6 +400,7 @@ "invalidJson": "配置必须是有效的 JSON" }, "messages": { + "idHelp": "由插件文件名派生的插件 ID。在 Agents 等配置选项中引用此插件时使用。", "configHelp": "使用键值对配置插件。如果插件无需配置则留空。", "configValidationError": "配置验证失败:", "schemaRenderError": "无法渲染配置表单。此插件的 schema 定义可能无效。", @@ -401,10 +409,10 @@ "allUsersHelp": "启用时,插件将可以访问所有用户,包括将来创建的。", "noUsers": "未选择用户", "permissionReason": "原因", - "usersRequired": "此插件需要访问用户信息。请选择允许此插件访问的用户, 或启用 '允许所有用户'。", + "usersRequired": "此插件需要访问用户信息。请选择允许此插件访问的用户,或启用 '允许所有用户'。", "allLibrariesHelp": "启用时,插件将可以访问所有媒体库,包括将来创建的。", "noLibraries": "未选择媒体库", - "librariesRequired": "此插件需要访问媒体库信息。请选择允许此插件访问的媒体库, 或启用 '允许所有媒体库'。", + "librariesRequired": "此插件需要访问媒体库信息。请选择允许此插件访问的媒体库,或启用 '允许所有媒体库'。", "allowWriteAccessHelp": "启用时,插件将可以修改媒体库目录中的文件。默认情况下,插件仅拥有只读权限。", "requiredHosts": "必需的主机" }, @@ -566,6 +574,7 @@ "coverRemoved": "封面已移除", "coverUploadError": "上传封面时出错", "coverRemoveError": "移除封面时出错", + "metadataRefreshStarted": "正在后台刷新元数据", "note": "注意", "transcodingDisabled": "出于安全原因,从 Web 界面更改转码配置的功能已被禁用。要更改(编辑或新增)转码选项,请在启用 %{config} 选项的情况下重新启动服务器。", "transcodingEnabled": "Navidrome 当前与 %{config} 一起使用,可以通过从 Web 界面配置转码选项来执行任意命令。建议禁用此选项,并且仅在需要配置转码选项时启用此功能。", @@ -598,7 +607,7 @@ "shareOriginalFormat": "分享原始格式", "shareDialogTitle": "分享 %{resource} '%{name}'", "shareBatchDialogTitle": "分享 %{smart_count} 项 %{resource}", - "shareCopyToClipboard": "复制到剪切板: Ctrl+C, Enter", + "shareCopyToClipboard": "复制到剪切板: Ctrl+C,Enter", "shareSuccess": "URL 已复制: %{url}", "shareFailure": "URL 复制失败: %{url}", "downloadDialogTitle": "下载 %{resource} '%{name}' (%{size})", diff --git a/resources/i18n/zh-Hant.json b/resources/i18n/zh-Hant.json index d00ae2ac3..7b8819da3 100644 --- a/resources/i18n/zh-Hant.json +++ b/resources/i18n/zh-Hant.json @@ -93,7 +93,8 @@ "addToPlaylist": "加入至播放清單", "download": "下載", "info": "取得資訊", - "share": "分享" + "share": "分享", + "refresh": "重新整理中繼資料" }, "lists": { "all": "所有", @@ -155,11 +156,13 @@ "newPassword": "新密碼", "token": "權杖", "lastAccessAt": "上次存取", - "libraries": "媒體庫" + "libraries": "媒體庫", + "scrobbleFilter": "紀錄過濾條件" }, "helperTexts": { "name": "您的名稱會在下次登入時生效", - "libraries": "為該使用者選擇指定媒體庫,留空則使用預設媒體庫" + "libraries": "為該使用者選擇指定媒體庫,留空則使用預設媒體庫", + "scrobbleFilter": "符合這些智慧播放清單規則的歌曲將不會傳送至 Last.fm、ListenBrainz 或 Scrobbler 插件。語法與行為皆與智慧播放清單的 JSON 格式相同。範例:{\"all\":[{\"lt\":{\"rating\":4}}]}。若留空則會記錄所有播放紀錄。本機播放次數不受影響。" }, "notifications": { "created": "使用者已建立", @@ -173,7 +176,8 @@ "adminAutoLibraries": "管理員預設可存取所有媒體庫" }, "validation": { - "librariesRequired": "非管理員使用者必須至少選擇一個媒體庫" + "librariesRequired": "非管理員使用者必須至少選擇一個媒體庫", + "invalidScrobbleFilter": "必須是有效的智慧播放清單規則。不支援數量限制、位移及重新整理延遲。" } }, "player": { @@ -196,6 +200,9 @@ "targetFormat": "目標格式", "defaultBitRate": "預設位元率", "command": "指令" + }, + "choices": { + "noDefaultBitRate": "無" } }, "playlist": { @@ -210,7 +217,8 @@ "songCount": "歌曲數", "comment": "註解", "sync": "自動匯入", - "path": "匯入來源" + "path": "匯入來源", + "starred": "收藏" }, "actions": { "selectPlaylist": "選取播放清單:", @@ -403,7 +411,8 @@ "requiredHosts": "必要的 Hosts", "configValidationError": "設定驗證失敗:", "schemaRenderError": "無法顯示設定表單。外掛的 schema 可能無效。", - "allowWriteAccessHelp": "啟用後,外掛可以修改媒體庫目錄中的檔案。 預設情況下,外掛具有唯讀權限。" + "allowWriteAccessHelp": "啟用後,外掛可以修改媒體庫目錄中的檔案。 預設情況下,外掛具有唯讀權限。", + "idHelp": "插件 ID,衍生自其檔案名稱。在 Agents 等設定選項中參照此插件時請使用此 ID。" }, "placeholders": { "configKey": "鍵", @@ -599,7 +608,8 @@ "coverUploaded": "已更新封面圖", "coverRemoved": "已移除封面圖", "coverUploadError": "上傳封面圖時發生錯誤", - "coverRemoveError": "移除封面圖時發生錯誤" + "coverRemoveError": "移除封面圖時發生錯誤", + "metadataRefreshStarted": "正在背景重新整理中繼資料" }, "menu": { "library": "媒體庫", @@ -634,7 +644,8 @@ "multipleLibraries": "已選 %{selected} 共 %{total} 媒體庫", "selectLibraries": "選取媒體庫", "none": "無" - } + }, + "onlyFavourites": "僅顯示收藏" }, "player": { "playListsText": "播放佇列", 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 diff --git a/scanner/controller.go b/scanner/controller.go index df5aeb6f9..1b13c1846 100644 --- a/scanner/controller.go +++ b/scanner/controller.go @@ -20,11 +20,12 @@ import ( "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server/events" "github.com/navidrome/navidrome/utils/pl" + "github.com/navidrome/navidrome/utils/singleton" "golang.org/x/time/rate" ) var ( - ErrAlreadyScanning = errors.New("already scanning") + ErrAlreadyScanning = model.ErrAlreadyScanning ) func New(rootCtx context.Context, ds model.DataStore, broker events.Broker, @@ -93,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 @@ -107,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, @@ -125,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) @@ -184,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) } @@ -237,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 @@ -303,14 +304,14 @@ func LockForMaintenance() (func(), bool) { return scanMaintenanceMux.Unlock, true } -// EffectiveFullScan reports whether a scan was requested as full or will resume an interrupted -// full scan in one of the included libraries. +// EffectiveFullScan reports whether a scan was requested as full, will resume an interrupted full scan, +// or will rescan a library in full because its PID config changed, in one of the included libraries. func EffectiveFullScan(ctx context.Context, ds model.DataStore, fullScan bool, targets []model.ScanTarget) bool { if fullScan { return true } return anyIncludedLibrary(ctx, ds, targets, func(library model.Library) bool { - return library.FullScanInProgress + return library.FullScanInProgress || library.NeedsPIDRescan() }) } @@ -323,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 } @@ -387,3 +388,12 @@ func (s *controller) trackProgress(ctx context.Context, progress <-chan *Progres func (s *controller) sendMessage(ctx context.Context, status *events.ScanStatus) { s.broker.SendBroadcastMessage(ctx, status) } + +// GetInstance returns the scanner singleton: Status reads the progress counters of the controller +// running the scan, and scheduler, watcher and signal scans do not start from the API's injector. +func GetInstance(rootCtx context.Context, ds model.DataStore, broker events.Broker, + pls playlists.Playlists, m metrics.Metrics) model.Scanner { + return singleton.GetInstance(func() *controller { + return New(rootCtx, ds, broker, pls, m).(*controller) + }) +} diff --git a/scanner/controller_test.go b/scanner/controller_test.go index 7ffdd69d4..974540e32 100644 --- a/scanner/controller_test.go +++ b/scanner/controller_test.go @@ -2,6 +2,7 @@ package scanner_test import ( "context" + "time" "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" @@ -35,7 +36,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 +44,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) @@ -70,14 +71,21 @@ var _ = Describe("EffectiveFullScan", func() { var ds *tests.MockDataStore BeforeEach(func() { + pid := model.Library{}.EffectivePID() libraries := &tests.MockLibraryRepo{} libraries.SetData(model.Libraries{ - {ID: 1, FullScanInProgress: true}, - {ID: 2}, + {ID: 1, FullScanInProgress: true, ScannedPIDAlbum: pid.Album, ScannedPIDTrack: pid.Track}, + {ID: 2, ScannedPIDAlbum: pid.Album, ScannedPIDTrack: pid.Track}, + {ID: 3, LastScanAt: time.Now(), PIDAlbum: "folder", ScannedPIDAlbum: pid.Album, ScannedPIDTrack: pid.Track}, }) ds = &tests.MockDataStore{MockedLibrary: libraries} }) + It("detects a library that needs a full rescan for a PID change", func() { + targets := []model.ScanTarget{{LibraryID: 3, FolderPath: "."}} + Expect(scanner.EffectiveFullScan(GinkgoT().Context(), ds, false, targets)).To(BeTrue()) + }) + It("detects an interrupted full scan in a targeted library", func() { targets := []model.ScanTarget{{LibraryID: 1, FolderPath: "."}} Expect(scanner.EffectiveFullScan(context.Background(), ds, false, targets)).To(BeTrue()) @@ -92,3 +100,13 @@ var _ = Describe("EffectiveFullScan", func() { Expect(scanner.EffectiveFullScan(context.Background(), ds, false, targets)).To(BeFalse()) }) }) + +var _ = Describe("GetInstance", func() { + It("returns the same controller to every caller", func() { + ds := &tests.MockDataStore{} + pls := playlists.NewPlaylists(ds, artwork.NewUploader(ds)) + a := scanner.GetInstance(context.Background(), ds, events.NoopBroker(), pls, metrics.NewNoopInstance()) + b := scanner.GetInstance(context.Background(), ds, events.NoopBroker(), pls, metrics.NewNoopInstance()) + Expect(a).To(BeIdenticalTo(b)) + }) +}) 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/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_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_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/scanner/phase_1_folders.go b/scanner/phase_1_folders.go index 6107b3316..4edaecacd 100644 --- a/scanner/phase_1_folders.go +++ b/scanner/phase_1_folders.go @@ -40,26 +40,30 @@ func createPhaseFolders(ctx context.Context, state *scanState, ds model.DataStor if err != nil { log.Error(ctx, "Scanner: Error creating scan context", "lib", lib.Name, err) state.sendError(err) + state.markFailed(lib.ID) continue } 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 { - lib model.Library - fs storage.MusicFS - lastUpdates map[string]model.FolderUpdateInfo // Holds last update info for all (DB) folders in this library - targetFolders []string // Specific folders to scan (including all descendants) - lock sync.Mutex - numFolders atomic.Int64 + lib model.Library + fs storage.MusicFS + lastUpdates map[string]model.FolderUpdateInfo // Holds last update info for all (DB) folders in this library + targetFolders []string // Specific folders to scan (including all descendants) + prevAlbumPIDConf string // Album PID spec of the last finished scan, only when it differs from the current one + lock sync.Mutex + numFolders atomic.Int64 } 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) } @@ -75,16 +79,32 @@ func newScanJob(ctx context.Context, ds model.DataStore, lib model.Library, full return nil, fmt.Errorf("getting fs for library: %w", err) } + pid := lib.EffectivePID() + if lib.NeedsPIDRescan() { + msg := "Scanner: PID config changed, rescanning library in full" + if len(targetFolders) > 0 { + msg = "Scanner: PID config changed, rescanning target folders in full" + } + log.Info(ctx, msg, "lib", lib.Name, "targetFolders", targetFolders, + "album", pid.Album, "track", pid.Track, "scannedAlbum", lib.ScannedPIDAlbum, "scannedTrack", lib.ScannedPIDTrack) + fullScan = true + } + var prevAlbumPIDConf string + if lib.ScannedPIDAlbum != pid.Album { + prevAlbumPIDConf = lib.ScannedPIDAlbum + } + // Ensure FullScanInProgress reflects the current scan request. // This is important when resuming an interrupted quick scan as a full scan: // the DB may have FullScanInProgress=false, but we need it true for isOutdated() to work correctly. lib.FullScanInProgress = lib.FullScanInProgress || fullScan return &scanJob{ - lib: lib, - fs: fsys, - lastUpdates: lastUpdates, - targetFolders: targetFolders, + lib: lib, + fs: fsys, + lastUpdates: lastUpdates, + targetFolders: targetFolders, + prevAlbumPIDConf: prevAlbumPIDConf, }, nil } @@ -120,12 +140,13 @@ func (j *scanJob) createFolderEntry(path string) *folderEntry { // The phaseFolders struct implements the phase interface, providing methods to produce // folder entries, process folders, persist changes to the database, and log the results. type phaseFolders struct { - jobs []*scanJob - ds model.DataStore - ctx context.Context - state *scanState - prevAlbumPIDConf string - imageChanges *imageChangeCollector + jobs []*scanJob + ds model.DataStore + 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 + imageChanges *imageChangeCollector } func (p *phaseFolders) description() string { @@ -134,25 +155,19 @@ 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, "") - if err != nil { - return fmt.Errorf("getting album PID conf: %w", err) - } - // TODO Parallelize multiple job when we have multiple libraries 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, @@ -169,7 +184,7 @@ func (p *phaseFolders) producer() ppl.Producer[*folderEntry] { // Check if folder is outdated if folder.isOutdated() { - if !p.state.fullScan { + if !folder.job.lib.FullScanInProgress { // Ancestor folders need a row even with no files of their own: artwork // resolution climbs them, and an image added later needs a state to diff. if folder.isEmpty() && folder.isNew() { @@ -208,9 +223,12 @@ 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{ + cursor, err := p.ds.MediaFile().GetCursor(p.ctx, model.QueryOptions{ Filters: squirrel.And{squirrel.Eq{"folder_id": entry.id}}, }) if err != nil { @@ -232,7 +250,7 @@ func (p *phaseFolders) processFolder(entry *folderEntry) (*folderEntry, error) { for afPath, af := range entry.audioFiles { fullPath := path.Join(entry.path, afPath) dbTrack, foundInDB := dbTracks[fullPath] - if !foundInDB || p.state.fullScan { + if !foundInDB || entry.job.lib.FullScanInProgress { filesToImport[fullPath] = dbTrack } else { info, err := af.Info() @@ -282,18 +300,18 @@ func (p *phaseFolders) loadTagsFromFiles(entry *folderEntry, toImport map[string } for filePath, info := range allInfo { md := metadata.New(filePath, info) - track := md.ToMediaFile(entry.job.lib.ID, entry.id) + track := md.ToMediaFile(entry.job.lib, entry.id) tracks = append(tracks, track) for _, t := range track.Tags.FlattenAll() { uniqueTags[t.ID] = t } // Keep track of any album ID changes, to reassign annotations later - prevAlbumID := "" + prevAlbumID := track.AlbumID if prev := toImport[filePath]; prev != nil { prevAlbumID = prev.AlbumID - } else { - prevAlbumID = md.AlbumID(track, p.prevAlbumPIDConf) + } else if entry.job.prevAlbumPIDConf != "" { + prevAlbumID = md.AlbumID(track, entry.job.prevAlbumPIDConf) } _, ok := entry.albumIDMap[track.AlbumID] if prevAlbumID != track.AlbumID && !ok { @@ -331,135 +349,136 @@ 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() + 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(ctx, folder) + if err != nil { + return fmt.Errorf("persisting folder: %w", err) + } + + // Save all tags to DB + 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(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(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) + } + 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(ctx, &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().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(ctx, 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(ctx, 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() + enqueue := queue.Enqueue + if entry.job.lib.FullScanInProgress { + enqueue = queue.EnqueueIfMissing + } + if err := enqueue(ctx, 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 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 == "" { @@ -468,13 +487,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) @@ -499,34 +518,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().MarkMissing(ctx, 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().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 - _, 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().TouchByMissingFolder(ctx); 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..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 @@ -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().Put(ctx, &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().Delete(ctx, 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().ReassignAnnotation(ctx, 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().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) + } + } + // 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().DeleteAllMissing(ctx) + 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 d54ceee40..f61aa4244 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" @@ -129,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}, @@ -142,16 +144,97 @@ 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)) }) + 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().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &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().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &matchedTrack) + + _, err := phase.processMissingTracks(&missingTracks{ + missing: []model.MediaFile{missingTrack}, + matched: []model.MediaFile{matchedTrack}, + }) + Expect(err).ToNot(HaveOccurred()) + + movedTrack, err := ds.MediaFile().Get(ctx, "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().Put(ctx, &missingTrack) + _ = ds.MediaFile().Put(ctx, &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} - _ = 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}, @@ -163,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)) }) @@ -172,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}, @@ -185,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)) }) @@ -195,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}, @@ -210,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)) }) @@ -220,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}, @@ -235,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)) @@ -250,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}, @@ -266,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)) }) @@ -278,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}, @@ -287,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()) @@ -311,7 +394,7 @@ var _ = Describe("phaseMissingTracks", func() { When("PurgeMissing is 'always'", func() { BeforeEach(func() { conf.Server.Scanner.PurgeMissing = consts.PurgeMissingAlways - mr.CountAllValue = 3 + mr.SetCountAll(3) mr.DeleteAllMissingValue = 3 }) It("should purge missing files", func() { @@ -325,7 +408,7 @@ var _ = Describe("phaseMissingTracks", func() { When("PurgeMissing is 'full'", func() { BeforeEach(func() { conf.Server.Scanner.PurgeMissing = consts.PurgeMissingFull - mr.CountAllValue = 2 + mr.SetCountAll(2) mr.DeleteAllMissingValue = 2 }) It("should not purge missing files if not a full scan", func() { @@ -346,7 +429,7 @@ var _ = Describe("phaseMissingTracks", func() { When("PurgeMissing is 'never'", func() { BeforeEach(func() { conf.Server.Scanner.PurgeMissing = consts.PurgeMissingNever - mr.CountAllValue = 1 + mr.SetCountAll(1) mr.DeleteAllMissingValue = 1 }) It("should not purge missing files", func() { @@ -431,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"}, @@ -446,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)) }) @@ -483,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"}, @@ -498,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)) }) @@ -529,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"}, @@ -587,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"}, @@ -603,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)) }) @@ -635,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"}, @@ -650,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)) }) @@ -705,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"}, @@ -721,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)) }) @@ -761,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) }) @@ -785,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}, @@ -796,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)) }) @@ -822,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)) }) @@ -856,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) @@ -881,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" @@ -916,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() { @@ -950,10 +1033,41 @@ 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)) }) }) }) }) + +// 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..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 @@ -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().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 { 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().RefreshPlayCounts(ctx) + 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().RefreshPlayCounts(ctx) + return txErr + }, "scanner: refresh artist play counts") if err != nil { return fmt.Errorf("refreshing artist annotations: %w", err) } 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 baa8b749a..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) @@ -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().Put(ctx, 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. "+ @@ -110,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 } @@ -147,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) @@ -164,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 d73007bdd..305c443f4 100644 --- a/scanner/scanner.go +++ b/scanner/scanner.go @@ -4,13 +4,11 @@ import ( "context" "fmt" "maps" - "path/filepath" "slices" "sync/atomic" "time" ppl "github.com/google/go-pipeline/pkg/pipeline" - "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/playlists" "github.com/navidrome/navidrome/log" @@ -32,6 +30,7 @@ type scanState struct { libraries model.Libraries // Store libraries list for consistency across phases targets map[int][]string // Optional: map[libraryID][]folderPaths for selective scans totalLibraryCount int // Total number of libraries (unfiltered), for cross-library move detection + failedLibs map[int]bool // Libraries that could not be scanned in this run } func (s *scanState) sendProgress(info *ProgressInfo) { @@ -48,29 +47,15 @@ func (s *scanState) sendWarning(msg string) { s.sendProgress(&ProgressInfo{Warning: msg}) } -func (s *scanState) sendError(err error) { - s.sendProgress(&ProgressInfo{Error: err.Error()}) +func (s *scanState) markFailed(libID int) { + if s.failedLibs == nil { + s.failedLibs = map[int]bool{} + } + s.failedLibs[libID] = true } -// libraryRelativePath rebases an absolute scan target path onto the library root, since the -// scanner's fs.FS only accepts paths relative to it. Relative paths, and absolute paths outside -// the library root, are returned unchanged. -func libraryRelativePath(libPath, folderPath string) string { - if !filepath.IsAbs(folderPath) { - return folderPath - } - // The library root may be relative (e.g. the default "./music"); it must be made absolute - // to match against an absolute target, and it resolves against the same cwd as the scanner's fs. - absLib, err := filepath.Abs(libPath) - if err != nil { - return folderPath - } - rel, err := filepath.Rel(absLib, folderPath) - if err != nil || !filepath.IsLocal(rel) { - return folderPath - } - // The scanner's fs.FS is an io/fs, which always uses forward slashes. - return filepath.ToSlash(rel) +func (s *scanState) sendError(err error) { + s.sendProgress(&ProgressInfo{Error: err.Error()}) } func (s *scannerImpl) scanFolders(ctx context.Context, fullScan bool, targets []model.ScanTarget, progress chan<- *ProgressInfo) { @@ -88,7 +73,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 @@ -104,7 +89,7 @@ func (s *scannerImpl) scanFolders(ctx context.Context, fullScan bool, targets [] }) for _, target := range targets { - folderPath := libraryRelativePath(libPaths[target.LibraryID], target.FolderPath) + folderPath := model.LibraryRelativePath(libPaths[target.LibraryID], target.FolderPath) if folderPath == "" { folderPath = "." } @@ -131,19 +116,23 @@ 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 { for _, lib := range state.libraries { + // A pending PID rescan already restarts in full through its own job + if lib.NeedsPIDRescan() { + continue + } if lib.FullScanInProgress { 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 +179,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}) @@ -215,9 +204,14 @@ func (s *scannerImpl) prepareLibrariesForScan(ctx context.Context, state *scanSt var successfulLibs []model.Library for _, lib := range state.libraries { - if lib.LastScanStartedAt.IsZero() { + // A library with a changed PID config restarts its scan: resuming would skip the folders that + // the interrupted scan already processed with the old config + pidRescan := lib.NeedsPIDRescan() + if lib.LastScanStartedAt.IsZero() || pidRescan { // 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().ScanBegin(ctx, lib.ID, state.fullScan || pidRescan) + }, "scanner: begin library scan") if err != nil { log.Error(ctx, "Scanner: Error marking scan start", "lib", lib.Name, err) state.sendWarning(err.Error()) @@ -225,7 +219,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()) @@ -253,7 +247,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 +258,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 +278,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().EnqueueAllMissing(ctx, 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) @@ -308,7 +304,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) @@ -316,7 +312,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().UpdateCounts(ctx) + }, "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 +327,22 @@ 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().ScanEnd(ctx, 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) - 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) - return fmt.Errorf("updating album PID conf: %w", err) + // A selective scan covers only part of the library, so the rest may still use the old PID + // config. A library that could not be scanned did not apply it either. + if !state.isSelectiveScan() && !state.failedLibs[lib.ID] { + if err := tx.Library().SetScannedPID(ctx, lib.ID, lib.EffectivePID()); err != nil { + return fmt.Errorf("updating PID conf for %s: %w", lib.Name, 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) + if err := tx.Library().RefreshStats(ctx, lib.ID); err != nil { + 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_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_internal_test.go b/scanner/scanner_internal_test.go index 0778bd6ec..e8abb7c7d 100644 --- a/scanner/scanner_internal_test.go +++ b/scanner/scanner_internal_test.go @@ -4,8 +4,6 @@ package scanner import ( "context" "errors" - "os" - "path/filepath" "sync/atomic" ppl "github.com/google/go-pipeline/pkg/pipeline" @@ -13,43 +11,6 @@ import ( . "github.com/onsi/gomega" ) -var _ = Describe("libraryRelativePath", func() { - // Paths are built with filepath so the "absolute" cases stay absolute on every OS - // (a Unix-style "/foo" is not absolute on Windows). - libRoot, _ := filepath.Abs(filepath.Join("jukebox", "collection")) - outside, _ := filepath.Abs(filepath.Join("somewhere", "else")) - - It("returns a relative path unchanged", func() { - Expect(libraryRelativePath(libRoot, "_Collection")).To(Equal("_Collection")) - }) - - It("rebases an absolute target when the library root is relative", func() { - cwd, err := os.Getwd() - Expect(err).ToNot(HaveOccurred()) - Expect(libraryRelativePath(filepath.Join("music", "library"), filepath.Join(cwd, "music", "library", "rock"))).To(Equal("rock")) - }) - - It("rebases an absolute path that equals the library root to '.'", func() { - Expect(libraryRelativePath(libRoot, libRoot)).To(Equal(".")) - }) - - It("rebases an absolute path under the library root", func() { - Expect(libraryRelativePath(libRoot, filepath.Join(libRoot, "_Collection"))).To(Equal("_Collection")) - }) - - It("handles a trailing slash on the library path", func() { - Expect(libraryRelativePath(libRoot+string(filepath.Separator), filepath.Join(libRoot, "_Collection"))).To(Equal("_Collection")) - }) - - It("leaves an absolute path outside the library root unchanged", func() { - Expect(libraryRelativePath(libRoot, outside)).To(Equal(outside)) - }) - - It("returns an empty path unchanged", func() { - Expect(libraryRelativePath(libRoot, "")).To(Equal("")) - }) -}) - type mockPhase struct { num int produceFunc func() ppl.Producer[int] diff --git a/scanner/scanner_multilibrary_test.go b/scanner/scanner_multilibrary_test.go index c0d5d4ece..f6634c875 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,17 +822,183 @@ 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()) }) }) + + Context("Per-library PID config", func() { + albumsOf := func(libID int) model.Albums { + // The mock datastore's GC is a no-op, so run the real one to purge the albums left + // empty by a regroup, as the scanner does in production + Expect(ds.RealDS.GC(ctx)).To(Succeed()) + albums, err := ds.Album().GetAll(ctx, model.QueryOptions{ + Filters: squirrel.Eq{"library_id": libID, "missing": false}, + Sort: "name", + }) + Expect(err).ToNot(HaveOccurred()) + return albums + } + trackByTitle := func(libID int, title string) model.MediaFile { + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{ + Filters: squirrel.Eq{"library_id": libID, "title": title}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(mfs).To(HaveLen(1)) + return mfs[0] + } + rockTitles := func() []string { + mfs, err := ds.MediaFile().GetAll(ctx, model.QueryOptions{Filters: squirrel.Eq{"library_id": lib1.ID}}) + Expect(err).ToNot(HaveOccurred()) + return slice.Map(mfs, func(mf model.MediaFile) string { return mf.Title }) + } + // changeRockInDB edits the rock track in the DB only. A full rescan of the rock library would + // restore the title from the file tags, a quick scan leaves it alone. + changeRockInDB := func() { + _, err := db.Db().ExecContext(ctx, "update media_file set title = 'Changed In DB' where library_id = ?", lib1.ID) + Expect(err).ToNot(HaveOccurred()) + } + // changeBlueTrainInDB does the same for one jazz track + changeBlueTrainInDB := func() { + _, err := db.Db().ExecContext(ctx, "update media_file set title = 'Blue Train In DB' where library_id = ? and title = 'Blue Train'", lib2.ID) + Expect(err).ToNot(HaveOccurred()) + } + + BeforeEach(func() { + beatles := template(_t{"albumartist": "The Beatles", "album": "Abbey Road", "year": 1969}) + _ = createFS("rock", fstest.MapFS{ + "The Beatles/Abbey Road/01 - Come Together.mp3": beatles(track(1, "Come Together")), + }) + + miles := template(_t{"albumartist": "Miles Davis", "album": "Kind of Blue", "year": 1959}) + coltrane := template(_t{"albumartist": "John Coltrane", "album": "Giant Steps", "year": 1960}) + blueTrain := template(_t{"albumartist": "John Coltrane", "album": "Blue Train", "year": 1957}) + _ = createFS("jazz", fstest.MapFS{ + "Loose/01 - So What.mp3": miles(track(1, "So What")), + "Loose/02 - Giant Steps.mp3": coltrane(track(1, "Giant Steps")), + "Coltrane/Blue Train/01 - Blue Train.mp3": blueTrain(track(1, "Blue Train")), + }) + }) + + It("regroups only the library whose PID config changed, keeping annotations", func() { + Expect(runScanner(ctx, true)).To(Succeed()) + Expect(albumsOf(lib2.ID)).To(HaveLen(3)) + + // Star Blue Train, to check the star follows the album to its new ID + oldBlueTrain := trackByTitle(lib2.ID, "Blue Train") + Expect(ds.Album().SetStar(ctx, true, oldBlueTrain.AlbumID)).To(Succeed()) + changeRockInDB() + + lib2.PIDAlbum = "folder" + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) + Expect(runScanner(ctx, false)).To(Succeed()) + + // Jazz is grouped by folder now: "Loose" is one album + Expect(albumsOf(lib2.ID)).To(HaveLen(2)) + Expect(trackByTitle(lib2.ID, "So What").AlbumID).To(Equal(trackByTitle(lib2.ID, "Giant Steps").AlbumID)) + + newBlueTrain := trackByTitle(lib2.ID, "Blue Train") + Expect(newBlueTrain.AlbumID).ToNot(Equal(oldBlueTrain.AlbumID)) + album, err := ds.Album().Get(ctx, newBlueTrain.AlbumID) + Expect(err).ToNot(HaveOccurred()) + Expect(album.Starred).To(BeTrue()) + + // Rock only got a quick scan + Expect(rockTitles()).To(ConsistOf("Changed In DB")) + + jazz, err := ds.Library().Get(ctx, lib2.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(jazz.ScannedPIDAlbum).To(Equal("folder")) + Expect(jazz.PIDChanged()).To(BeFalse()) + rock, err := ds.Library().Get(ctx, lib1.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(rock.PIDChanged()).To(BeFalse()) + }) + + It("rescans only libraries that follow the global config", func() { + lib2.PIDAlbum = "folder" + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) + Expect(runScanner(ctx, true)).To(Succeed()) + changeRockInDB() + changeBlueTrainInDB() + + conf.Server.PID.Album = "album" + Expect(runScanner(ctx, false)).To(Succeed()) + + // Rock follows the global config, so it was rescanned in full and its title restored + Expect(rockTitles()).To(ConsistOf("Come Together")) + // Jazz has its own override, so it only got a quick scan + trackByTitle(lib2.ID, "Blue Train In DB") + jazz, err := ds.Library().Get(ctx, lib2.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(jazz.ScannedPIDAlbum).To(Equal("folder")) + }) + + It("restarts an interrupted scan when the PID config changed meanwhile", func() { + Expect(runScanner(ctx, true)).To(Succeed()) + + // Simulate a quick scan of jazz that was interrupted after it had processed every folder: + // the folders were updated after the (old) scan start time + _, err := db.Db().ExecContext(ctx, "update library set last_scan_started_at = ?, full_scan_in_progress = false where id = ?", + time.Now().Add(-time.Hour), lib2.ID) + Expect(err).ToNot(HaveOccurred()) + + lib2.PIDAlbum = "folder" + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) + Expect(runScanner(ctx, false)).To(Succeed()) + + // Every folder was revisited with the new config + Expect(albumsOf(lib2.ID)).To(HaveLen(2)) + }) + + It("does not turn an interrupted PID rescan into a full scan of every library", func() { + Expect(runScanner(ctx, true)).To(Succeed()) + changeRockInDB() + lib2.PIDAlbum = "folder" + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) + + // Simulate a PID full scan of jazz that was interrupted + Expect(ds.Library().ScanBegin(ctx, lib2.ID, true)).To(Succeed()) + Expect(runScanner(ctx, false)).To(Succeed()) + + // Rock only got a quick scan, jazz was rescanned with the new config + Expect(rockTitles()).To(ConsistOf("Changed In DB")) + Expect(albumsOf(lib2.ID)).To(HaveLen(2)) + }) + + It("does not record the PID config for a library that could not be scanned", func() { + Expect(runScanner(ctx, true)).To(Succeed()) + broken := model.Library{Name: "Broken", Path: "unregistered:///music", PIDAlbum: "folder"} + Expect(ds.Library().Put(ctx, &broken)).To(Succeed()) + + // The scan reports an error for the broken library, and still finishes the others + _ = runScanner(ctx, false) + + reloaded, err := ds.Library().Get(ctx, broken.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(reloaded.PIDChanged()).To(BeTrue()) + }) + + It("does not record the PID config after a selective scan", func() { + Expect(runScanner(ctx, true)).To(Succeed()) + lib2.PIDAlbum = "folder" + Expect(ds.Library().Put(ctx, &lib2)).To(Succeed()) + + _, err := s.ScanFolders(ctx, false, []model.ScanTarget{{LibraryID: lib2.ID, FolderPath: "Loose"}}) + Expect(err).ToNot(HaveOccurred()) + + jazz, err := ds.Library().Get(ctx, lib2.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(jazz.PIDChanged()).To(BeTrue()) + }) + }) }) 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 8542b3ac6..4ce8cce1e 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:" @@ -71,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 @@ -83,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 { @@ -101,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 } @@ -133,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), @@ -143,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), @@ -156,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!")), @@ -170,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 { @@ -197,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()) }) @@ -206,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"})) }) @@ -227,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"), @@ -241,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()) @@ -250,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)) @@ -259,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()) @@ -272,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)) }) }) @@ -284,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 } @@ -457,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 { @@ -479,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 { @@ -505,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( @@ -544,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) }) @@ -552,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")) @@ -562,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)")) }) @@ -571,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)) }) @@ -584,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") @@ -605,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") @@ -634,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))) @@ -643,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()) @@ -664,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") @@ -691,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") @@ -705,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") @@ -730,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") @@ -744,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))) @@ -783,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))) }) @@ -808,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 }) } @@ -853,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()) @@ -880,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()) @@ -908,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()) @@ -936,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()) @@ -963,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()) @@ -985,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()) @@ -1013,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)) @@ -1028,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)) @@ -1049,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()) @@ -1064,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)) @@ -1083,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()) @@ -1098,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()) @@ -1118,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()) @@ -1137,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()) @@ -1156,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, @@ -1202,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()) @@ -1221,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()) @@ -1240,11 +1247,58 @@ 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().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()) + }) + }) }) +// 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}) + list, err := ds.MediaFile().FindByPaths(ctx, []string{path}) if err != nil { return nil, err } @@ -1258,13 +1312,19 @@ func createFindByPath(ctx context.Context, ds model.DataStore) func(string) (*mo type mockMediaFileRepo struct { model.MediaFileRepository GetMissingAndMatchingError error + cursorCalls atomic.Int32 } -func (m *mockMediaFileRepo) GetMissingAndMatching(libId int) (model.MediaFileCursor, error) { +func (m *mockMediaFileRepo) GetCursor(ctx context.Context, options ...model.QueryOptions) (model.MediaFileCursor, error) { + m.cursorCalls.Add(1) + return m.MediaFileRepository.GetCursor(ctx, options...) +} + +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 { @@ -1272,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/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/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/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/auth.go b/server/auth.go index 37a318a83..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 @@ -96,6 +96,16 @@ func buildAuthPayload(user *model.User) map[string]any { return payload } +// MaxLoginBodySize bounds the payload of unauthenticated login routes across all APIs. +const MaxLoginBodySize = 8 << 10 + +func LimitLoginBody(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + r.Body = http.MaxBytesReader(w, r.Body, MaxLoginBodySize) + next.ServeHTTP(w, r) + }) +} + func getCredentialsFromBody(r *http.Request) (username string, password string, err error) { data := make(map[string]string) decoder := json.NewDecoder(r.Body) @@ -117,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 @@ -147,15 +157,16 @@ 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, err) + log.Error(ctx, "Could not create initial user", "user", initialUser.UserName, err) + return fmt.Errorf("creating initial user: %w", err) } 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 } @@ -165,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 } @@ -207,14 +218,14 @@ func UsernameFromExtAuthHeader(r *http.Request) string { log.Error("ExtAuth enabled but no proxy IP found in request context. Please report this error.") return "" } - if !validateIPAgainstList(reverseProxyIp, conf.Server.ExtAuth.TrustedSources) { - log.Warn(r.Context(), "IP is not whitelisted for external authentication", "proxy-ip", reverseProxyIp, "client-ip", r.RemoteAddr) - return "" - } username := r.Header.Get(conf.Server.ExtAuth.UserHeader) if username == "" { return "" } + if !validateIPAgainstList(reverseProxyIp, conf.Server.ExtAuth.TrustedSources) { + log.Warn(r.Context(), "IP is not whitelisted for external authentication", "proxy-ip", reverseProxyIp, "client-ip", r.RemoteAddr) + return "" + } log.Trace(r, "Found username in ExtAuth.UserHeader", "username", username) return username } @@ -233,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) @@ -298,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 } @@ -366,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{ @@ -382,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 e021c82a8..1095fafc9 100644 --- a/server/auth_test.go +++ b/server/auth_test.go @@ -4,6 +4,7 @@ import ( "context" "crypto/md5" "encoding/json" + "errors" "fmt" "net/http" "net/http/httptest" @@ -15,15 +16,24 @@ import ( "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/id" "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" + "github.com/sirupsen/logrus" + "github.com/sirupsen/logrus/hooks/test" ) 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 @@ -44,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()) @@ -64,6 +74,26 @@ var _ = Describe("Auth", func() { }) }) + Describe("createAdminUser", func() { + It("returns the error when the user cannot be saved", func() { + ds = &tests.MockDataStore{MockedUser: &tests.MockedUserRepo{Error: errors.New("db is down")}} + err := createAdminUser(context.Background(), ds, "johndoe", "secret") + Expect(err).To(MatchError(ContainSubstring("db is down"))) + }) + }) + + 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" @@ -75,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() @@ -200,9 +230,16 @@ var _ = Describe("Auth", func() { Expect(resp.Code).To(Equal(http.StatusUnauthorized)) }) + It("rejects a request body larger than the limit", func() { + body := `{"username":"janedoe", "password":"abc123", "padding":"` + strings.Repeat("x", MaxLoginBodySize) + `"}` + req = httptest.NewRequest("POST", "/login", strings.NewReader(body)) + LimitLoginBody(http.HandlerFunc(login(ds))).ServeHTTP(resp, req) + Expect(resp.Code).To(Equal(http.StatusUnprocessableEntity)) + }) + 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)) @@ -218,6 +255,56 @@ var _ = Describe("Auth", func() { }) }) + Describe("UsernameFromExtAuthHeader", func() { + var hook *test.Hook + var r *http.Request + + BeforeEach(func() { + conf.Server.ExtAuth.TrustedSources = "192.168.0.0/16" + prevLevel := log.CurrentLevel() + l, h := test.NewNullLogger() + hook = h + prevLogger := log.SetDefaultLogger(l) + log.SetLevel(log.LevelWarn) + DeferCleanup(func() { + log.SetDefaultLogger(prevLogger) + log.SetLevel(prevLevel) + }) + r = httptest.NewRequest("GET", "/", nil) + }) + + warnings := func() []*logrus.Entry { + var ws []*logrus.Entry + for _, e := range hook.AllEntries() { + if e.Level == logrus.WarnLevel { + ws = append(ws, e) + } + } + return ws + } + + It("returns the username from a trusted source", func() { + r.Header.Set("Remote-User", "janedoe") + r = r.WithContext(request.WithReverseProxyIp(r.Context(), "192.168.0.42")) + Expect(UsernameFromExtAuthHeader(r)).To(Equal("janedoe")) + Expect(warnings()).To(BeEmpty()) + }) + + It("does not warn when an untrusted source sends no user header", func() { + r = r.WithContext(request.WithReverseProxyIp(r.Context(), "8.8.8.8")) + Expect(UsernameFromExtAuthHeader(r)).To(BeEmpty()) + Expect(warnings()).To(BeEmpty()) + }) + + It("warns when an untrusted source sends the user header", func() { + r.Header.Set("Remote-User", "janedoe") + r = r.WithContext(request.WithReverseProxyIp(r.Context(), "8.8.8.8")) + Expect(UsernameFromExtAuthHeader(r)).To(BeEmpty()) + Expect(warnings()).To(HaveLen(1)) + Expect(warnings()[0].Message).To(Equal("IP is not whitelisted for external authentication")) + }) + }) + Describe("tokenFromHeader", func() { It("returns the token when the Authorization header is set correctly", func() { req := httptest.NewRequest("GET", "/", nil) @@ -316,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", @@ -338,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()) }) @@ -353,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/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/server/initial_setup.go b/server/initial_setup.go index 7e974dc21..462e22e54 100644 --- a/server/initial_setup.go +++ b/server/initial_setup.go @@ -16,34 +16,37 @@ import ( func initialSetup(ds model.DataStore) { ctx := context.TODO() - _ = ds.WithTx(func(tx model.DataStore) error { - if err := tx.Library(ctx).StoreMusicFolder(); err != nil { + err := ds.WithTx(func(tx model.DataStore) error { + 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 { + log.Fatal("Error running initial setup", err) + } } // 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 { - 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(ctx, &initialUser); err != nil { + return fmt.Errorf("could not create initial admin user: %w", err) } } - return err + return nil } func checkFFmpegInstallation() { @@ -70,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/server/initial_setup_test.go b/server/initial_setup_test.go index 982046f78..0c85d9d0a 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,28 +10,52 @@ import ( . "github.com/onsi/gomega" ) +type failingPutUserRepo struct { + model.UserRepository + err error +} + +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}} +} + 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(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(ctx, ds, "pass123")).To(MatchError(boom)) }) }) }) diff --git a/server/jellyfin/README.md b/server/jellyfin/README.md index 5b9a2dba4..9315ac4f9 100644 --- a/server/jellyfin/README.md +++ b/server/jellyfin/README.md @@ -22,6 +22,10 @@ Enabled = true ServerName = "My Music Server" # Optional: usernames to show in the client login user-picker (default: none). See "Public user list". ExposedPublicUsers = "alice, bob" +# Optional: answer LAN auto-discovery broadcasts on UDP 7359 (default: false). See "Auto discovery". +AutoDiscovery = true +# Optional: let users sign in new devices with a 6-digit code (default: true). See "Quick Connect". +QuickConnect = false # Optional: max collection responses streaming at once (default: half the DB connection pool, # min 2). Each streaming response holds a DB connection for its whole duration; excess requests # queue rather than fail. @@ -34,6 +38,9 @@ or via environment variables: ND_JELLYFIN_ENABLED=true 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: @@ -46,6 +53,26 @@ All the paths below are relative to that base URL (e.g. `System/Info/Public` mea `http://localhost:4533/jellyfin/System/Info/Public`). Routes are matched **case-insensitively**, since real Jellyfin clients (and `jellyfin-apiclient-python`) send mixed-case paths. +## Auto discovery + +With `AutoDiscovery = true`, Navidrome answers the Jellyfin LAN discovery broadcast +(`who is JellyfinServer?` on UDP port 7359), so clients list the server without a typed URL. +It is off by default because a real Jellyfin server on the same host owns that port. If the port +is taken, Navidrome logs a warning and keeps running without discovery. + +The advertised address is `BaseURL` when it includes a host. Otherwise it is the bind `Address` +when that is a specific IP, or else the local IP that faces the requesting client, plus `Port`. With a +unix socket `Address` there is no port to advertise, so discovery only starts when `BaseURL` has a host. + +Discovery answers on all IPv4 interfaces, but the advertised address follows `BaseURL`, `Address` and +`Port`. If `Address` is a loopback or a single interface IP and `BaseURL` has no host, clients on other +networks get an address they cannot reach. Set `BaseURL` to the address clients should use. + +Docker: use host networking (`network_mode: host` / `--network host`). On Linux, bridge mode does not +deliver broadcasts to the container, even with `-p 7359:7359/udp`, so clients never find the server. +With host networking the server sees the host's IP, so `BaseURL` is not needed for discovery. If host +networking is not an option, leave discovery off. Keep UDP 7359 on the LAN: never forward it from the internet. + ## Authentication Jellyfin clients authenticate with `POST /Users/AuthenticateByName` using the user's Navidrome @@ -60,6 +87,27 @@ surface. Access tokens do not expire, matching real Jellyfin. They are revoked by a password change, which bumps the user's token epoch. +### Quick Connect + +Quick Connect signs a new device in without typing a password. The client shows a 6-digit code, +a signed-in user approves it, and the client gets its `AccessToken`. It is on by default +(`Jellyfin.QuickConnect`); when off, `GET QuickConnect/Enabled` returns `false` and every other +Quick Connect call returns 401. + +1. `POST QuickConnect/Initiate` (public; needs `Client`, `Device`, `DeviceId` and `Version` in the + auth header) returns the `Code` and a `Secret`. +2. The user approves the code, either in the Navidrome web UI (user menu → **Quick Connect**, which + shows the app and device before approving) or from a signed-in Jellyfin client with + `POST QuickConnect/Authorize?Code=`. Admins may pass `UserId` to approve for another user. +3. The client polls `GET QuickConnect/Connect?Secret=` until `Authenticated` is `true`. +4. `POST Users/AuthenticateWithQuickConnect` with `{"Secret": "..."}` returns the same result as + `AuthenticateByName`. + +Pending codes live in memory and expire after 10 minutes (a server restart drops them). Unlike +Jellyfin, a secret signs in only once. `Initiate`, `AuthenticateWithQuickConnect` and code approval +(`Authorize` and the web UI) are rate-limited per IP like the login; `Connect` is not, since some +clients poll it every second. + ### Public user list (login picker) `GET /Users/Public` lets a client render a login user-picker (tap a user, then just type the @@ -93,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 @@ -105,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 @@ -118,18 +168,19 @@ returns direct children only (no tracks — no track is a library's direct child | Area | Endpoints | |---|---| -| Handshake / system | `GET System/Info/Public`, `GET System/Info` (authenticated), `GET`/`POST System/Ping`, `GET System/Endpoint` (authenticated), `GET QuickConnect/Enabled` | +| Handshake / system | `GET System/Info/Public`, `GET System/Info` (authenticated), `GET`/`POST System/Ping`, `GET System/Endpoint` (authenticated) | +| 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/{itemId}/InstantMix` | +| 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` | +| 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) | | Lyrics | `GET Audio/{itemId}/Lyrics` | -| Playback reporting | `POST Sessions/Playing`, `POST Sessions/Playing/Progress`, `POST Sessions/Playing/Stopped`, `POST Sessions/Capabilities[/Full]` | -| Playlists | `POST Playlists`, `GET Playlists/{playlistId}`, `POST Playlists/{playlistId}` (rename / visibility / replace tracks), `GET Playlists/{playlistId}/Items`, `POST`/`DELETE Playlists/{playlistId}/Items`, `GET Playlists/{playlistId}/Users[/{userId}]` | +| Playback reporting | `POST Sessions/Playing`, `POST Sessions/Playing/Progress`, `POST Sessions/Playing/Stopped`, `POST Sessions/Playing/Ping` (no-op), `POST Sessions/Capabilities[/Full]` | +| Playlists | `POST Playlists`, `GET Playlists/{playlistId}`, `POST Playlists/{playlistId}` (rename / visibility / replace tracks), `GET Playlists/{playlistId}/Items`, `POST`/`DELETE Playlists/{playlistId}/Items` (`POST` honors `position`), `POST Playlists/{playlistId}/Items/{entryId}/Move/{newIndex}`, `GET Playlists/{playlistId}/Users[/{userId}]` | | Real-time | `GET socket` (WebSocket; keeps clients like Finamp from 404-loop-reconnecting) | | AudioMuse-AI (see below) | `GET AudioMuseAI/info`, `GET AudioMuseAI/health`, `GET AudioMuseAI/similar_tracks`, `GET AudioMuseAI/find_path` | @@ -219,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) @@ -308,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/... @@ -317,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 — the reason `jellyfinVersion` is 10.9.11. - 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/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.go b/server/jellyfin/api.go index 6484a3bb4..4a471d1f0 100644 --- a/server/jellyfin/api.go +++ b/server/jellyfin/api.go @@ -3,11 +3,9 @@ package jellyfin import ( "encoding/json" "net/http" - "sync" "time" "github.com/go-chi/chi/v5" - "github.com/go-chi/httprate" "golang.org/x/sync/singleflight" "github.com/navidrome/navidrome/conf" @@ -16,6 +14,7 @@ import ( "github.com/navidrome/navidrome/core/external" "github.com/navidrome/navidrome/core/lyrics" "github.com/navidrome/navidrome/core/playlists" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/core/scrobbler" "github.com/navidrome/navidrome/core/sonic" "github.com/navidrome/navidrome/core/stream" @@ -42,18 +41,18 @@ type Router struct { broker events.Broker lyricsCache cache.SimpleCache[string, model.LyricList] similarFlight singleflight.Group - serverIDMu sync.Mutex + quickConnect quickconnect.QuickConnect serverIDVal string } func New(ds model.DataStore, artwork artwork.Artwork, streamer stream.MediaStreamer, transcodeDecider stream.TranscodeDecider, players core.Players, scrobbler scrobbler.PlayTracker, playlists playlists.Playlists, provider external.Provider, - sonicSvc sonic.Engine, lyricsSvc lyrics.Lyrics, broker events.Broker) *Router { + sonicSvc sonic.Engine, lyricsSvc lyrics.Lyrics, broker events.Broker, quickConnect quickconnect.QuickConnect) *Router { r := &Router{ ds: ds, artwork: artwork, streamer: streamer, transcodeDecider: transcodeDecider, players: players, scrobbler: scrobbler, playlists: playlists, provider: provider, - sonic: sonicSvc, lyrics: lyricsSvc, broker: broker, + sonic: sonicSvc, lyrics: lyricsSvc, broker: broker, quickConnect: quickConnect, lyricsCache: cache.NewSimpleCache[string, model.LyricList](cache.Options{ SizeLimit: 1000, DefaultTTL: 5 * time.Minute, @@ -78,13 +77,17 @@ func (api *Router) routes() http.Handler { inner.Post("/system/ping", api.ping) inner.Get("/quickconnect/enabled", api.quickConnectEnabled) // Rate-limit the password login, mirroring the native /auth/login: it's an unauthenticated - // brute-force surface, so it must share the same per-IP throttle when one is configured. + // brute-force surface, so it must share the same per-client throttle when one is configured. + login := inner.With(server.LimitLoginBody) if conf.Server.AuthRequestLimit > 0 { - limiter := httprate.LimitByIP(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength) - inner.With(limiter).Post("/users/authenticatebyname", api.authenticateByName) - } else { - inner.Post("/users/authenticatebyname", api.authenticateByName) + login = login.With(server.ClientIPRateLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength)) } + login.Post("/users/authenticatebyname", api.authenticateByName) + quickConnectLogin := login.With(requireQuickConnect) + quickConnectLogin.Post("/quickconnect/initiate", api.quickConnectInitiate) + quickConnectLogin.Post("/users/authenticatewithquickconnect", api.authenticateWithQuickConnect) + // Not rate-limited: Finamp and Streamyfin poll it every second while the code is shown. + inner.With(requireQuickConnect).Get("/quickconnect/connect", api.quickConnectConnect) inner.Get("/users/public", api.getPublicUsers) // Images are intentionally public: artwork isn't sensitive, matching Jellyfin's image handling. @@ -95,6 +98,8 @@ func (api *Router) routes() http.Handler { conf.Server.DevArtworkThrottleBacklogTimeout)) r.Get("/items/{itemId}/images/{type}", api.getItemImage) r.Get("/items/{itemId}/images/{type}/{index}", api.getItemImage) + r.Head("/items/{itemId}/images/{type}", api.getItemImage) + r.Head("/items/{itemId}/images/{type}/{index}", api.getItemImage) }) inner.Group(func(r chi.Router) { @@ -109,6 +114,12 @@ func (api *Router) routes() http.Handler { r.Get("/users/{userId}/views", api.getUserViews) r.Get("/users/me", api.getCurrentUser) r.Get("/users/{userId}", api.getCurrentUser) + // Throttled like login so a signed-in user cannot enumerate other people's pending codes. + approve := r.With(requireQuickConnect) + if conf.Server.AuthRequestLimit > 0 { + approve = approve.With(server.ClientIPRateLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength)) + } + approve.Post("/quickconnect/authorize", api.quickConnectAuthorize) // Cursor-backed collections: each streams straight from the DB, holding a connection for the // whole client-paced response, so enough slow clients would take the entire pool and stall the @@ -147,6 +158,12 @@ func (api *Router) routes() http.Handler { r.Get("/items/{itemId}/similar", api.getSimilarItems) r.Get("/albums/{itemId}/similar", api.getSimilarAlbums) r.Get("/items/{itemId}/instantmix", api.getInstantMix) + r.Get("/songs/{itemId}/instantmix", api.getInstantMix) + r.Get("/albums/{itemId}/instantmix", api.getInstantMix) + r.Get("/artists/{itemId}/instantmix", api.getInstantMix) + r.Get("/playlists/{itemId}/instantmix", api.getInstantMix) + r.Get("/artists/instantmix", api.getInstantMixByQuery) + r.Get("/musicgenres/instantmix", api.getInstantMixByQuery) r.Get("/genres", api.getGenres) r.Get("/musicgenres", api.getGenres) r.Get("/studios", api.getStudios) @@ -157,6 +174,7 @@ func (api *Router) routes() http.Handler { r.Post("/playlists/{playlistId}", api.updatePlaylist) r.Post("/playlists/{playlistId}/items", api.addToPlaylist) r.Delete("/playlists/{playlistId}/items", api.removeFromPlaylist) + r.Post("/playlists/{playlistId}/items/{entryId}/move/{newIndex}", api.movePlaylistItem) r.Get("/playlists/{playlistId}/users", api.getPlaylistUsers) r.Get("/playlists/{playlistId}/users/{userId}", api.getPlaylistUser) @@ -167,7 +185,11 @@ func (api *Router) routes() http.Handler { r.Get("/audio/{itemId}/stream", api.streamAudio) r.Get("/audio/{itemId}/stream.{container}", api.streamAudio) - r.Get("/audio/{itemId}/universal", api.streamAudio) + r.Get("/audio/{itemId}/universal", api.streamUniversal) + // Fintunes probes these with HEAD for the content type before playing or downloading. + r.Head("/audio/{itemId}/stream", api.streamAudio) + r.Head("/audio/{itemId}/stream.{container}", api.streamAudio) + r.Head("/audio/{itemId}/universal", api.streamUniversal) r.Get("/audio/{itemId}/main.m3u8", api.streamHls) r.Get("/items/{itemId}/playbackinfo", api.getPlaybackInfo) r.Post("/items/{itemId}/playbackinfo", api.getPlaybackInfo) @@ -176,12 +198,15 @@ func (api *Router) routes() http.Handler { // /Audio/{id}/stream; /Download reuses the direct-play handler as Jellyfin serves the same file. r.Get("/items/{itemId}/file", api.streamFile) r.Get("/items/{itemId}/download", api.streamFile) + r.Head("/items/{itemId}/file", api.streamFile) + r.Head("/items/{itemId}/download", api.streamFile) r.Post("/sessions/playing", api.reportPlaybackStart) r.Post("/sessions/playing/progress", api.reportPlaybackProgress) r.Post("/sessions/playing/stopped", api.reportPlaybackStopped) - r.Post("/sessions/capabilities", api.postCapabilities) - r.Post("/sessions/capabilities/full", api.postCapabilities) + r.Post("/sessions/playing/ping", api.acknowledge) + r.Post("/sessions/capabilities", api.acknowledge) + r.Post("/sessions/capabilities/full", api.acknowledge) // Real-time clients (e.g. Finamp) open this right after login; without it they 404-loop-reconnect. r.Get("/socket", api.handleSocket) @@ -215,8 +240,7 @@ func (api *Router) ok(w http.ResponseWriter, r *http.Request, payload any) { api.writeItems(w, r, materialized(p)) return case dto.BaseItemDto: - p.ServerId = api.serverID(r.Context()) - payload = p + payload = stampItem(p, api.serverID(r.Context()), requestFields(r)) } w.Header().Set("Content-Type", "application/json; charset=utf-8") if err := json.NewEncoder(w).Encode(payload); err != nil { diff --git a/server/jellyfin/api_test.go b/server/jellyfin/api_test.go index fe69c8321..0c6ac9b98 100644 --- a/server/jellyfin/api_test.go +++ b/server/jellyfin/api_test.go @@ -1,14 +1,17 @@ package jellyfin import ( + "context" "net/http" "net/http/httptest" "strings" "time" + "github.com/go-chi/chi/v5/middleware" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" @@ -16,9 +19,15 @@ 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) + api := New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/System/Info/Public", nil) api.ServeHTTP(w, r) @@ -26,7 +35,7 @@ var _ = Describe("Router", func() { }) It("returns 404 JSON for unknown routes", func() { - api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Nonexistent/Route", nil) api.ServeHTTP(w, r) @@ -36,7 +45,7 @@ var _ = Describe("Router", func() { }) It("returns 404 JSON for a known path with an unsupported method", func() { - api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) w := httptest.NewRecorder() r := httptest.NewRequest("PATCH", "/System/Info/Public", nil) api.ServeHTTP(w, r) @@ -47,13 +56,13 @@ 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()) fp := &fakePlayers{} - api := New(ds, nil, nil, nil, fp, nil, nil, nil, nil, nil, nil) + api := New(ds, nil, nil, nil, fp, nil, nil, nil, nil, nil, nil, nil) w := httptest.NewRecorder() r := httptest.NewRequest("GET", "/Users/Me", nil) @@ -70,7 +79,7 @@ var _ = Describe("Router", func() { DeferCleanup(configtest.SetupConfig()) conf.Server.AuthRequestLimit = 2 conf.Server.AuthWindowLength = time.Minute - api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) login := func() int { w := httptest.NewRecorder() @@ -84,4 +93,53 @@ var _ = Describe("Router", func() { Expect(login()).To(Equal(http.StatusUnauthorized)) Expect(login()).To(Equal(http.StatusTooManyRequests)) }) + + It("rate-limits Quick Connect approval by IP when a login limit is configured", func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.AuthRequestLimit = 2 + conf.Server.AuthWindowLength = time.Minute + conf.Server.Jellyfin.QuickConnect = true + ds := &tests.MockDataStore{} + auth.Init(ds) + usr := model.User{ID: testID("alice"), UserName: "alice"} + 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()) + + authorize := func() int { + w := httptest.NewRecorder() + r := httptest.NewRequest("POST", "/QuickConnect/Authorize?code=000000", nil) + r.RemoteAddr = "10.0.0.1:1234" + r.Header.Set("X-Emby-Token", token) + api.ServeHTTP(w, r) + return w.Code + } + // An unknown code is 404; the limiter cuts in on the 3rd attempt with 429. + Expect(authorize()).To(Equal(http.StatusNotFound)) + Expect(authorize()).To(Equal(http.StatusNotFound)) + Expect(authorize()).To(Equal(http.StatusTooManyRequests)) + }) + + It("rate-limits AuthenticateByName by resolved client IP, not by the proxy connection", func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.AuthRequestLimit = 1 + conf.Server.AuthWindowLength = time.Minute + api := New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + // Every request arrives on the same proxy connection, so only the resolved client IP + // can separate the buckets. + handler := middleware.ClientIPFromHeader("X-Real-IP")(api) + + login := func(clientIP string) int { + w := httptest.NewRecorder() + r := httptest.NewRequest("POST", "/Users/AuthenticateByName", strings.NewReader(`{"Username":"x","Pw":"y"}`)) + r.RemoteAddr = "10.0.0.1:1234" + r.Header.Set("X-Real-IP", clientIP) + handler.ServeHTTP(w, r) + return w.Code + } + Expect(login("203.0.113.1")).To(Equal(http.StatusUnauthorized)) + Expect(login("203.0.113.1")).To(Equal(http.StatusTooManyRequests)) + Expect(login("203.0.113.2")).To(Equal(http.StatusUnauthorized)) + }) }) diff --git a/server/jellyfin/auth.go b/server/jellyfin/auth.go index e7070d341..264d5b6ca 100644 --- a/server/jellyfin/auth.go +++ b/server/jellyfin/auth.go @@ -24,16 +24,21 @@ 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) return } + api.signIn(w, r, usr) +} + +func (api *Router) signIn(w http.ResponseWriter, r *http.Request, usr *model.User) { + 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 { - log.Error(ctx, "Jellyfin API: could not update last login date", "username", body.Username, err) + 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) } token, err := auth.CreateAPIToken(usr, auth.AudienceJellyfin) @@ -42,12 +47,13 @@ func (api *Router) authenticateByName(w http.ResponseWriter, r *http.Request) { return } - // SessionInfo is omitted, not partially filled: a stub {Id, UserId} could fail a strict client's - // parse, and Finamp's login doesn't need it (its AuthenticationResult.sessionInfo is nullable). + a := parseMediaBrowserAuth(r) + serverID := api.serverID(ctx) api.ok(w, r, dto.AuthenticationResult{ - User: userToDto(usr, api.serverName(), api.serverID(ctx)), + User: userToDto(usr, serverName(), serverID), + SessionInfo: dto.NewSessionInfo(usr, a.Client, a.DeviceId, a.Device, a.Version, serverID), AccessToken: token, - ServerId: api.serverID(ctx), + ServerId: serverID, }) } diff --git a/server/jellyfin/auth_test.go b/server/jellyfin/auth_test.go index 5a613f828..dafad9244 100644 --- a/server/jellyfin/auth_test.go +++ b/server/jellyfin/auth_test.go @@ -9,6 +9,7 @@ import ( "github.com/navidrome/navidrome/core/auth" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/server" "github.com/navidrome/navidrome/server/jellyfin/dto" "github.com/navidrome/navidrome/tests" . "github.com/onsi/ginkgo/v2" @@ -16,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} }) @@ -47,10 +50,54 @@ var _ = Describe("AuthenticateByName", func() { Expect(res.User.Policy.EnableAllFolders).To(BeTrue()) Expect(res.User.Policy.EnableMediaPlayback).To(BeTrue()) Expect(res.User.Configuration).ToNot(BeNil()) + }) - // Ours is a partial SessionInfo; a strict client may fail to parse it, and Finamp's - // login doesn't require it, so it should be omitted entirely rather than sent partial. - Expect(res.SessionInfo).To(BeNil()) + Describe("SessionInfo", func() { + login := func(authHeader string) map[string]any { + w := httptest.NewRecorder() + r := httptest.NewRequest("POST", "/Users/AuthenticateByName", + strings.NewReader(`{"Username":"alice","Pw":"secret"}`)) + r.Header.Set("Authorization", authHeader) + api.authenticateByName(w, r) + Expect(w.Code).To(Equal(http.StatusOK)) + var raw map[string]any + Expect(json.Unmarshal(w.Body.Bytes(), &raw)).To(Succeed()) + Expect(raw).To(HaveKey("SessionInfo")) + return raw["SessionInfo"].(map[string]any) + } + const jellybox = `MediaBrowser Client="JellyBox", Device="Mac", DeviceId="dev-1", Version="2.1"` + + It("has the fields strict clients require", func() { + s := login(jellybox) + Expect(s["Id"]).To(And(BeAssignableToTypeOf(""), Not(BeEmpty()))) + Expect(s["UserId"]).To(Equal(dto.EncodeID(testID("u1")))) + Expect(s["LastActivityDate"]).To(And(BeAssignableToTypeOf(""), Not(BeEmpty()))) + for _, k := range []string{"SupportsRemoteControl", "SupportsMediaControl", "HasCustomDeviceName"} { + Expect(s[k]).To(BeAssignableToTypeOf(false), k) + } + ps, ok := s["PlayState"].(map[string]any) + Expect(ok).To(BeTrue()) + for _, k := range []string{"CanSeek", "IsPaused", "IsMuted"} { + Expect(ps[k]).To(BeAssignableToTypeOf(false), k) + } + }) + + It("describes the calling client", func() { + s := login(jellybox) + Expect(s).To(HaveKeyWithValue("UserName", "alice")) + Expect(s).To(HaveKeyWithValue("Client", "JellyBox")) + Expect(s).To(HaveKeyWithValue("DeviceName", "Mac")) + Expect(s).To(HaveKeyWithValue("DeviceId", "dev-1")) + Expect(s).To(HaveKeyWithValue("ApplicationVersion", "2.1")) + Expect(s).To(HaveKeyWithValue("IsActive", true)) + }) + + It("keeps the same Id for the same device, and a new one for another device", func() { + first := login(jellybox)["Id"] + Expect(login(jellybox)["Id"]).To(Equal(first)) + other := login(`MediaBrowser Client="JellyBox", Device="Mac", DeviceId="dev-2", Version="2.1"`) + Expect(other["Id"]).ToNot(Equal(first)) + }) }) It("records the login time, like the web UI login does", func() { @@ -60,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", @@ -91,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", @@ -101,3 +148,15 @@ var _ = Describe("AuthenticateByName", func() { Expect(w.Code).To(Equal(http.StatusUnauthorized)) }) }) + +var _ = Describe("AuthenticateByName body limit", func() { + It("rejects a request body larger than the limit", func() { + ds := &tests.MockDataStore{} + api := &Router{ds: ds} + w := httptest.NewRecorder() + body := `{"Username":"alice","Pw":"secret","Padding":"` + strings.Repeat("x", server.MaxLoginBodySize) + `"}` + r := httptest.NewRequest("POST", "/Users/AuthenticateByName", strings.NewReader(body)) + server.LimitLoginBody(http.HandlerFunc(api.authenticateByName)).ServeHTTP(w, r) + Expect(w.Code).To(Equal(http.StatusBadRequest)) + }) +}) 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/discovery.go b/server/jellyfin/discovery.go new file mode 100644 index 000000000..d111287a6 --- /dev/null +++ b/server/jellyfin/discovery.go @@ -0,0 +1,110 @@ +package jellyfin + +import ( + "context" + "encoding/json" + "net" + "strconv" + "strings" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/core/publicurl" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/utils/gg" +) + +// Jellyfin clients broadcast this query text to this UDP port. +const ( + discoveryPort = 7359 + discoveryQuery = "who is jellyfinserver?" +) + +type discoveryInfo struct { + Address string `json:"Address"` + Id string `json:"Id"` + Name string `json:"Name"` + EndpointAddress *string `json:"EndpointAddress"` +} + +// Discovery answers LAN auto-discovery broadcasts with the same identity the Router reports. +type Discovery struct { + ds model.DataStore + serverIDVal string +} + +func NewDiscovery(ds model.DataStore) *Discovery { + return &Discovery{ds: ds} +} + +func (d *Discovery) serverID(ctx context.Context) string { + return resolveServerID(ctx, d.ds, &d.serverIDVal) +} + +// Serve runs until ctx is done. A failed bind is only logged: discovery is best-effort. +func (d *Discovery) Serve(ctx context.Context) { + if !hasAdvertisableAddress() { + log.Warn(ctx, "Jellyfin API: auto-discovery is off, a unix socket server needs a BaseURL with a host to advertise") + return + } + // udp4 only: a dual-stack bind can share the port with another server and never get a packet. + conn, err := net.ListenPacket("udp4", net.JoinHostPort("0.0.0.0", strconv.Itoa(discoveryPort))) + if err != nil { + log.Warn(ctx, "Jellyfin API: auto-discovery is off, the UDP port is unavailable. Is another Jellyfin server running?", "port", discoveryPort, err) + return + } + log.Info(ctx, "Jellyfin API: listening for auto-discovery broadcasts", "port", discoveryPort) + d.ServeOn(ctx, conn) +} + +// ServeOn answers discovery queries on conn until ctx is done, then closes conn. +func (d *Discovery) ServeOn(ctx context.Context, conn net.PacketConn) { + defer conn.Close() + stop := context.AfterFunc(ctx, func() { _ = conn.Close() }) + defer stop() + buf := make([]byte, 1024) + for { + n, remote, err := conn.ReadFrom(buf) + if err != nil { + if ctx.Err() == nil { + log.Error(ctx, "Jellyfin API: auto-discovery listener stopped", err) + } + return + } + if !strings.Contains(strings.ToLower(string(buf[:n])), discoveryQuery) { + continue + } + info := discoveryInfo{Address: discoveryAddress(ctx, remote), Id: d.serverID(ctx), Name: serverName()} + res, _ := json.Marshal(info) + log.Debug(ctx, "Jellyfin API: answering auto-discovery request", "from", remote.String(), "address", info.Address) + if _, err := conn.WriteTo(res, remote); err != nil { + log.Debug(ctx, "Jellyfin API: could not answer auto-discovery request", "to", remote.String(), err) + } + } +} + +// Behind a unix socket nothing listens on Port, so only a BaseURL host gives clients an address. +func hasAdvertisableAddress() bool { + return conf.Server.BaseHost != "" || !strings.HasPrefix(conf.Server.Address, "unix:") +} + +func discoveryAddress(ctx context.Context, remote net.Addr) string { + scheme := gg.If(conf.Server.TLSEnabled(), "https", "http") + host := net.JoinHostPort(localIPFor(remote), strconv.Itoa(conf.Server.Port)) + return publicurl.AbsoluteURL(request.WithServerAddress(ctx, scheme, host), consts.URLPathJellyfinAPI, nil) +} + +// On a multi-homed host, only the interface that routes to the requester is reachable by it. +func localIPFor(remote net.Addr) string { + if ip := parseIP(conf.Server.Address); ip.IsValid() && !ip.IsUnspecified() { + return ip.String() + } + c, err := net.Dial("udp", remote.String()) + if err != nil { + return conf.Server.Address + } + defer c.Close() + return c.LocalAddr().(*net.UDPAddr).IP.String() +} diff --git a/server/jellyfin/discovery_test.go b/server/jellyfin/discovery_test.go new file mode 100644 index 000000000..bcd11fe76 --- /dev/null +++ b/server/jellyfin/discovery_test.go @@ -0,0 +1,175 @@ +package jellyfin + +import ( + "context" + "encoding/json" + "errors" + "net" + "os" + "time" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type failingConn struct { + net.PacketConn + closed bool +} + +func (c *failingConn) ReadFrom([]byte) (int, net.Addr, error) { + return 0, nil, errors.New("read failed") +} + +func (c *failingConn) Close() error { + c.closed = true + return nil +} + +var _ = Describe("Discovery", func() { + var d *Discovery + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.Jellyfin.ServerName = "Test Server" + conf.Server.Address = "0.0.0.0" + conf.Server.Port = 4533 + conf.Server.BaseHost = "" + conf.Server.BaseScheme = "" + conf.Server.BasePath = "" + conf.Server.TLSCert = "" + conf.Server.TLSKey = "" + d = &Discovery{} + }) + + DescribeTable("discoveryAddress", + func(setup func(), expected string) { + setup() + remote := &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 50000} + Expect(discoveryAddress(context.Background(), remote)).To(Equal(expected)) + }, + Entry("uses the BaseURL host and scheme when set", func() { + conf.Server.BaseScheme = "https" + conf.Server.BaseHost = "music.example.com" + conf.Server.BasePath = "/nd" + }, "https://music.example.com/nd/jellyfin"), + Entry("uses a specific bind Address with the Port", func() { + conf.Server.Address = "192.168.1.10" + }, "http://192.168.1.10:4533/jellyfin"), + Entry("falls back to the interface facing the requester when Address is unspecified", func() {}, + "http://127.0.0.1:4533/jellyfin"), + Entry("falls back to the interface facing the requester when Address is empty", func() { + conf.Server.Address = "" + }, "http://127.0.0.1:4533/jellyfin"), + Entry("advertises https when TLS is configured", func() { + conf.Server.TLSCert = "/path/cert.pem" + conf.Server.TLSKey = "/path/key.pem" + }, "https://127.0.0.1:4533/jellyfin"), + Entry("advertises http when only the TLS cert is configured", func() { + conf.Server.TLSCert = "/path/cert.pem" + }, "http://127.0.0.1:4533/jellyfin"), + Entry("keeps a path-only BaseURL as the path prefix", func() { + conf.Server.BasePath = "/music" + }, "http://127.0.0.1:4533/music/jellyfin"), + Entry("does not double the slash when BasePath has a trailing slash", func() { + conf.Server.BasePath = "/music/" + }, "http://127.0.0.1:4533/music/jellyfin"), + ) + + DescribeTable("hasAdvertisableAddress", + func(address, baseHost string, expected bool) { + conf.Server.Address = address + conf.Server.BaseHost = baseHost + Expect(hasAdvertisableAddress()).To(Equal(expected)) + }, + Entry("TCP listener", "0.0.0.0", "", true), + Entry("unix socket without a BaseURL host", "unix:/tmp/navidrome.sock", "", false), + Entry("unix socket behind a proxy named by BaseURL", "unix:/tmp/navidrome.sock", "music.example.com", true), + ) + + It("closes the connection when the read loop fails", func() { + fake := &failingConn{} + d.ServeOn(context.Background(), fake) + Expect(fake.closed).To(BeTrue()) + }) + + Describe("ServeOn", func() { + var ( + server net.PacketConn + client net.PacketConn + cancel context.CancelFunc + done chan struct{} + ) + + BeforeEach(func() { + var err error + server, err = net.ListenPacket("udp4", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + client, err = net.ListenPacket("udp4", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(client.Close) + + var ctx context.Context + ctx, cancel = context.WithCancel(context.Background()) + done = make(chan struct{}) + go func() { + defer close(done) + d.ServeOn(ctx, server) + }() + DeferCleanup(func() { + cancel() + Eventually(done).Should(BeClosed()) + }) + }) + + send := func(msg string) { + _, err := client.WriteTo([]byte(msg), server.LocalAddr()) + Expect(err).ToNot(HaveOccurred()) + } + receive := func() ([]byte, error) { + Expect(client.SetReadDeadline(time.Now().Add(time.Second))).To(Succeed()) + buf := make([]byte, 1024) + n, _, err := client.ReadFrom(buf) + return buf[:n], err + } + + It("answers the discovery query with the server identity", func() { + send("who is JellyfinServer?") + res, err := receive() + Expect(err).ToNot(HaveOccurred()) + + var info discoveryInfo + Expect(json.Unmarshal(res, &info)).To(Succeed()) + Expect(info.Address).To(Equal("http://127.0.0.1:4533/jellyfin")) + Expect(info.Id).To(Equal(d.serverID(context.Background()))) + Expect(info.Name).To(Equal("Test Server")) + Expect(string(res)).To(ContainSubstring(`"EndpointAddress":null`)) + }) + + It("matches the query case-insensitively", func() { + send("WHO IS JELLYFINSERVER?") + _, err := receive() + Expect(err).ToNot(HaveOccurred()) + }) + + // The loop is serial, so a reply to the first packet would arrive before the second's. + It("ignores unrelated packets", func() { + send("who is PlexServer?") + send("who is JellyfinServer?") + _, err := receive() + Expect(err).ToNot(HaveOccurred()) + Expect(client.SetReadDeadline(time.Now().Add(50 * time.Millisecond))).To(Succeed()) + _, _, err = client.ReadFrom(make([]byte, 1024)) + Expect(errors.Is(err, os.ErrDeadlineExceeded)).To(BeTrue()) + }) + + It("stops and closes the socket when the context is cancelled", func() { + cancel() + Eventually(done).Should(BeClosed()) + _, _, err := server.ReadFrom(make([]byte, 1)) + Expect(err).To(HaveOccurred()) + }) + }) +}) diff --git a/server/jellyfin/dto/dto.go b/server/jellyfin/dto/dto.go index 58a6b919e..4a28eb313 100644 --- a/server/jellyfin/dto/dto.go +++ b/server/jellyfin/dto/dto.go @@ -49,20 +49,23 @@ type BaseItemDto struct { // PlaylistItemId identifies an entry within a playlist listing (GET /Playlists/{id}/Items), // distinct from Id so a song appearing more than once can be removed by occurrence // (DELETE .../Items?EntryIds=...) rather than by song id. - PlaylistItemId string `json:"PlaylistItemId,omitempty"` - Type string `json:"Type"` - IsFolder bool `json:"IsFolder"` - MediaType string `json:"MediaType,omitempty"` - CollectionType string `json:"CollectionType,omitempty"` - LocationType string `json:"LocationType,omitempty"` - HasLyrics bool `json:"HasLyrics,omitempty"` - SortName string `json:"SortName,omitempty"` - Path string `json:"Path,omitempty"` - ParentId string `json:"ParentId,omitempty"` - RunTimeTicks int64 `json:"RunTimeTicks,omitempty"` - IndexNumber *int `json:"IndexNumber,omitempty"` - ParentIndexNumber *int `json:"ParentIndexNumber,omitempty"` - ProductionYear *int `json:"ProductionYear,omitempty"` + PlaylistItemId string `json:"PlaylistItemId,omitempty"` + Type string `json:"Type"` + IsFolder bool `json:"IsFolder"` + MediaType string `json:"MediaType,omitempty"` + CollectionType string `json:"CollectionType,omitempty"` + LocationType string `json:"LocationType,omitempty"` + HasLyrics *bool `json:"HasLyrics,omitempty"` + // ChannelId is always null for music, but Jellyfin emits it on every item and clients may require it. + ChannelId *string `json:"ChannelId"` + Tags []string `json:"Tags,omitzero"` + SortName string `json:"SortName,omitempty"` + Path string `json:"Path,omitempty"` + ParentId string `json:"ParentId,omitempty"` + RunTimeTicks int64 `json:"RunTimeTicks,omitempty"` + IndexNumber *int `json:"IndexNumber,omitempty"` + ParentIndexNumber *int `json:"ParentIndexNumber,omitempty"` + ProductionYear *int `json:"ProductionYear,omitempty"` // PremiereDate is the ISO 8601 release date; Finamp sorts "Latest Releases" by it client-side. PremiereDate *string `json:"PremiereDate,omitempty"` // DateCreated is the ISO 8601 date the item was added to the library; clients show it as @@ -73,17 +76,17 @@ type BaseItemDto struct { AlbumArtist string `json:"AlbumArtist,omitempty"` AlbumArtists []NameGuidPair `json:"AlbumArtists,omitempty"` AlbumPrimaryImageTag string `json:"AlbumPrimaryImageTag,omitempty"` - Artists []string `json:"Artists,omitempty"` + Artists []string `json:"Artists,omitzero"` ArtistItems []NameGuidPair `json:"ArtistItems,omitempty"` - Genres []string `json:"Genres,omitempty"` - GenreItems []NameGuidPair `json:"GenreItems,omitempty"` + Genres []string `json:"Genres,omitzero"` + GenreItems []NameGuidPair `json:"GenreItems,omitzero"` Studios []NameGuidPair `json:"Studios,omitempty"` NormalizationGain *float64 `json:"NormalizationGain,omitempty"` AlbumNormalizationGain *float64 `json:"AlbumNormalizationGain,omitempty"` ChildCount *int `json:"ChildCount,omitempty"` SongCount *int `json:"SongCount,omitempty"` AlbumCount *int `json:"AlbumCount,omitempty"` - ImageTags map[string]string `json:"ImageTags,omitempty"` + ImageTags map[string]string `json:"ImageTags"` // ImageBlurHashes is keyed by image type (e.g. "Primary") then image tag. Finamp uses it as a // de-dup key for image downloads (and a placeholder); absent, it warns the server isn't // calculating blurhashes. @@ -199,14 +202,51 @@ type UserConfiguration struct { CastReceiverId string `json:"CastReceiverId"` } +// SessionInfo mirrors real Jellyfin's SessionInfoDto. JellyBox requires Id and PlayState, and Finamp +// requires UserId, LastActivityDate and the bools, so none of those may be omitted. type SessionInfo struct { - Id string `json:"Id"` - UserId string `json:"UserId"` + Id string `json:"Id"` + UserId string `json:"UserId"` + UserName string `json:"UserName"` + Client string `json:"Client"` + DeviceId string `json:"DeviceId"` + DeviceName string `json:"DeviceName"` + ApplicationVersion string `json:"ApplicationVersion"` + ServerId string `json:"ServerId"` + LastActivityDate string `json:"LastActivityDate"` + IsActive bool `json:"IsActive"` + SupportsMediaControl bool `json:"SupportsMediaControl"` + SupportsRemoteControl bool `json:"SupportsRemoteControl"` + HasCustomDeviceName bool `json:"HasCustomDeviceName"` + PlayableMediaTypes []string `json:"PlayableMediaTypes"` + SupportedCommands []string `json:"SupportedCommands"` + AdditionalUsers []any `json:"AdditionalUsers"` + NowPlayingQueue []any `json:"NowPlayingQueue"` + PlayState PlayerStateInfo `json:"PlayState"` +} + +type PlayerStateInfo struct { + CanSeek bool `json:"CanSeek"` + IsPaused bool `json:"IsPaused"` + IsMuted bool `json:"IsMuted"` + RepeatMode string `json:"RepeatMode"` + PlaybackOrder string `json:"PlaybackOrder"` +} + +type QuickConnectResult struct { + Authenticated bool `json:"Authenticated"` + Secret string `json:"Secret"` + Code string `json:"Code"` + DeviceId string `json:"DeviceId"` + DeviceName string `json:"DeviceName"` + AppName string `json:"AppName"` + AppVersion string `json:"AppVersion"` + DateAdded string `json:"DateAdded"` } type AuthenticationResult struct { User *UserDto `json:"User"` - SessionInfo *SessionInfo `json:"SessionInfo,omitempty"` + SessionInfo *SessionInfo `json:"SessionInfo"` AccessToken string `json:"AccessToken"` ServerId string `json:"ServerId"` } diff --git a/server/jellyfin/dto/mappers.go b/server/jellyfin/dto/mappers.go index 1ebf66ca1..fc7e329de 100644 --- a/server/jellyfin/dto/mappers.go +++ b/server/jellyfin/dto/mappers.go @@ -7,6 +7,7 @@ import ( "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/id" "github.com/navidrome/navidrome/utils/slice" ) @@ -46,17 +47,24 @@ func premiereDate(date string, year int) *string { } d = fmt.Sprintf("%04d-01-01", year) } - s := d + "T00:00:00Z" + parsed, err := time.Parse(time.DateOnly, d) + if err != nil { + return nil + } + s := JellyfinDate(&parsed) return &s } -// jellyfinDate formats t as the ISO 8601 string clients expect, or "" for the zero time so the +// Dates use .NET's round-trip layout, 7 fractional digits and all: Manet rejects plain RFC3339. +const jellyfinDateLayout = "2006-01-02T15:04:05.0000000Z07:00" + +// JellyfinDate formats t as the date string clients expect, or "" for the zero time so the // field is omitted rather than sent as a meaningless epoch. -func jellyfinDate(t *time.Time) string { +func JellyfinDate(t *time.Time) string { if t == nil || t.IsZero() { return "" } - return t.UTC().Format(time.RFC3339) + return t.UTC().Format(jellyfinDateLayout) } // channelLayout maps a channel count to the label Jellyfin clients expect on a MediaStream. @@ -127,8 +135,7 @@ func UserData(a model.Annotations, itemID string) *UserItemDataDto { r := float64(a.Rating) * 2 // Navidrome 0-5 -> Jellyfin 0-10 d.Rating = &r } - if a.PlayDate != nil { - s := a.PlayDate.UTC().Format(time.RFC3339) + if s := JellyfinDate(a.PlayDate); s != "" { d.LastPlayedDate = &s } return d @@ -146,13 +153,13 @@ func SongToBaseItem(mf model.MediaFile, fields Fields) BaseItemDto { MediaType: "Audio", IsFolder: false, LocationType: "FileSystem", - HasLyrics: mf.HasEmbeddedLyrics(), + HasLyrics: new(mf.HasEmbeddedLyrics()), ParentId: albumID, Album: mf.Album, AlbumId: albumID, AlbumArtist: mf.AlbumArtist, RunTimeTicks: TicksFromSeconds(mf.Duration), - DateCreated: jellyfinDate(&mf.CreatedAt), + DateCreated: JellyfinDate(&mf.CreatedAt), Container: mf.Suffix, CanDownload: true, BackdropImageTags: []string{}, @@ -251,13 +258,14 @@ func AlbumToBaseItem(al model.Album, fields Fields) BaseItemDto { Id: EncodeID(al.ID), Type: "MusicAlbum", IsFolder: true, + LocationType: "FileSystem", ParentId: EncodeID(al.AlbumArtistID), AlbumArtist: al.AlbumArtist, Album: al.Name, ChildCount: new(al.SongCount), SongCount: new(al.SongCount), RunTimeTicks: TicksFromSeconds(al.Duration), - DateCreated: jellyfinDate(&al.CreatedAt), + DateCreated: JellyfinDate(&al.CreatedAt), ImageBlurHashes: blurs, PrimaryImageAspectRatio: ratio, BackdropImageTags: []string{}, @@ -266,6 +274,10 @@ func AlbumToBaseItem(al model.Album, fields Fields) BaseItemDto { if tag != "" { item.ImageTags = map[string]string{"Primary": tag} } + item.Artists = []string{} + if al.AlbumArtist != "" { + item.Artists = append(item.Artists, al.AlbumArtist) + } if al.AlbumArtistID != "" { item.AlbumArtists = []NameGuidPair{{Name: al.AlbumArtist, Id: EncodeID(al.AlbumArtistID)}} item.ArtistItems = item.AlbumArtists @@ -306,7 +318,7 @@ func ArtistToBaseItem(ar model.Artist, fields Fields) BaseItemDto { IsFolder: true, AlbumCount: new(ar.AlbumCount), SongCount: new(ar.SongCount), - DateCreated: jellyfinDate(ar.CreatedAt), + DateCreated: JellyfinDate(ar.CreatedAt), ImageBlurHashes: blurs, PrimaryImageAspectRatio: ratio, BackdropImageTags: []string{}, @@ -321,6 +333,26 @@ func ArtistToBaseItem(ar model.Artist, fields Fields) BaseItemDto { return item } +// LibraryToBaseItem maps a library to the CollectionFolder item clients browse as a top-level node. +// Manet keeps no library, and so syncs nothing, unless it carries the fields Jellyfin sends here. +func LibraryToBaseItem(lib model.Library) BaseItemDto { + id := EncodeLibraryID(lib.ID) + return BaseItemDto{ + Id: id, + Name: lib.Name, + SortName: lib.Name, + Type: "CollectionFolder", + CollectionType: "music", + IsFolder: true, + Path: lib.Path, + LocationType: "FileSystem", + DateCreated: JellyfinDate(&lib.CreatedAt), + ChildCount: new(lib.TotalAlbums), + UserData: &UserItemDataDto{Key: id, ItemId: id}, + BackdropImageTags: []string{}, + } +} + func GenreToBaseItem(g model.Genre) BaseItemDto { return BaseItemDto{ Name: g.Name, @@ -354,6 +386,7 @@ func PlaylistToBaseItem(p model.Playlist, fields Fields) BaseItemDto { MediaType: "Audio", ChildCount: new(p.SongCount), RunTimeTicks: TicksFromSeconds(p.Duration), + DateCreated: JellyfinDate(&p.CreatedAt), ImageBlurHashes: blurs, PrimaryImageAspectRatio: ratio, BackdropImageTags: []string{}, @@ -362,6 +395,10 @@ func PlaylistToBaseItem(p model.Playlist, fields Fields) BaseItemDto { if tag != "" { item.ImageTags = map[string]string{"Primary": tag} } + // Playlists have no sort tag; the repository orders them by name. + if fields.Has("SortName") { + item.SortName = p.Name + } return item } @@ -410,3 +447,25 @@ func LyricDtoFromLyrics(mf model.MediaFile, lyrics model.Lyrics) LyricDto { } return d } + +// NewSessionInfo's Id is stable per client install, as Jellyfin reuses a device's session across logins. +func NewSessionInfo(u *model.User, client, deviceID, deviceName, version, serverID string) *SessionInfo { + now := time.Now() + return &SessionInfo{ + Id: EncodeID(id.NewHash(client, deviceID)), + UserId: EncodeID(u.ID), + UserName: u.UserName, + Client: client, + DeviceId: deviceID, + DeviceName: deviceName, + ApplicationVersion: version, + ServerId: serverID, + LastActivityDate: JellyfinDate(&now), + IsActive: true, + PlayableMediaTypes: []string{"Audio"}, + SupportedCommands: []string{}, + AdditionalUsers: []any{}, + NowPlayingQueue: []any{}, + PlayState: PlayerStateInfo{RepeatMode: "RepeatNone", PlaybackOrder: "Default"}, + } +} diff --git a/server/jellyfin/dto/mappers_test.go b/server/jellyfin/dto/mappers_test.go index 576482e61..259904673 100644 --- a/server/jellyfin/dto/mappers_test.go +++ b/server/jellyfin/dto/mappers_test.go @@ -17,10 +17,10 @@ var _ = Describe("mappers", func() { ID: testID("song-1"), Title: "Song", Album: "Alb", AlbumID: testID("alb-1"), Artist: "Art", AlbumArtist: "AA", TrackNumber: 3, DiscNumber: 1, Year: 1999, Duration: 60, Size: 2_500_000, - Genres: []model.Genre{{ID: testID("1"), Name: "genre 1"}, {ID: testID("2"), Name: "genre 2"}}, + Genres: []model.Genre{{ID: testID("1"), Name: "genre 1"}, {ID: testID("2"), Name: "genre 2"}}, + PlayCount: 2, + Starred: true, } - mf.PlayCount = 2 - mf.Starred = true item := SongToBaseItem(mf, nil) Expect(item.Type).To(Equal("Audio")) Expect(item.MediaType).To(Equal("Audio")) @@ -104,10 +104,10 @@ var _ = Describe("mappers", func() { }) It("sets HasLyrics from the media file's lyrics", func() { - Expect(SongToBaseItem(mf, nil).HasLyrics).To(BeTrue()) - Expect(SongToBaseItem(model.MediaFile{ID: testID("s2"), Title: "No Lyrics"}, nil).HasLyrics).To(BeFalse()) + Expect(*SongToBaseItem(mf, nil).HasLyrics).To(BeTrue()) + Expect(*SongToBaseItem(model.MediaFile{ID: testID("s2"), Title: "No Lyrics"}, nil).HasLyrics).To(BeFalse()) // "[]" is the no-lyrics sentinel, not a truthy value. - Expect(SongToBaseItem(model.MediaFile{ID: testID("s3"), Title: "Empty Lyrics", Lyrics: "[]"}, nil).HasLyrics).To(BeFalse()) + Expect(*SongToBaseItem(model.MediaFile{ID: testID("s3"), Title: "Empty Lyrics", Lyrics: "[]"}, nil).HasLyrics).To(BeFalse()) }) }) @@ -120,7 +120,7 @@ var _ = Describe("mappers", func() { It("sets DateCreated from the media file's CreatedAt", func() { mf := model.MediaFile{ID: testID("s1"), Title: "Song", CreatedAt: time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC)} - Expect(SongToBaseItem(mf, nil).DateCreated).To(Equal("2024-01-15T10:30:00Z")) + Expect(SongToBaseItem(mf, nil).DateCreated).To(Equal("2024-01-15T10:30:00.0000000Z")) }) It("omits DateCreated when CreatedAt is the zero time", func() { @@ -282,7 +282,7 @@ var _ = Describe("mappers", func() { mf.PlayDate = &playDate item := SongToBaseItem(mf, nil) Expect(item.UserData.LastPlayedDate).NotTo(BeNil()) - Expect(*item.UserData.LastPlayedDate).To(Equal(playDate.Format(time.RFC3339))) + Expect(*item.UserData.LastPlayedDate).To(Equal("2023-05-17T12:30:00.0000000Z")) }) It("maps an album to a MusicAlbum folder item", func() { @@ -299,6 +299,9 @@ var _ = Describe("mappers", func() { Expect(*item.ChildCount).To(Equal(10)) Expect(item.ImageTags).To(HaveKeyWithValue("Primary", testID("alb-1"))) Expect(item.ImageBlurHashes).To(BeNil()) + // Jellyfin sends both on every album; Manet stops its sync on an album without them. + Expect(item.Artists).To(Equal([]string{"AA"})) + Expect(item.LocationType).To(Equal("FileSystem")) Expect(item.Genres).To(Equal([]string{"genre 1", "genre 2"})) Expect(item.GenreItems).To(Equal([]NameGuidPair{{Id: EncodeID(testID("1")), Name: "genre 1"}, {Id: EncodeID(testID("2")), Name: "genre 2"}})) }) @@ -426,22 +429,22 @@ var _ = Describe("mappers", func() { It("serializes a full date", func() { mf := model.MediaFile{ID: testID("s1"), Title: "Song", Date: "2007-02-01", Year: 2007} item := SongToBaseItem(mf, nil) - Expect(*item.PremiereDate).To(Equal("2007-02-01T00:00:00Z")) + Expect(*item.PremiereDate).To(Equal("2007-02-01T00:00:00.0000000Z")) }) It("pads a year-only date so clients can parse it", func() { mf := model.MediaFile{ID: testID("s1"), Title: "Song", Date: "2007", Year: 2007} - Expect(*SongToBaseItem(mf, nil).PremiereDate).To(Equal("2007-01-01T00:00:00Z")) + Expect(*SongToBaseItem(mf, nil).PremiereDate).To(Equal("2007-01-01T00:00:00.0000000Z")) }) It("pads a year-month date", func() { mf := model.MediaFile{ID: testID("s1"), Title: "Song", Date: "2007-02"} - Expect(*SongToBaseItem(mf, nil).PremiereDate).To(Equal("2007-02-01T00:00:00Z")) + Expect(*SongToBaseItem(mf, nil).PremiereDate).To(Equal("2007-02-01T00:00:00.0000000Z")) }) It("falls back to the year when no date tag exists", func() { mf := model.MediaFile{ID: testID("s1"), Title: "Song", Year: 1999} - Expect(*SongToBaseItem(mf, nil).PremiereDate).To(Equal("1999-01-01T00:00:00Z")) + Expect(*SongToBaseItem(mf, nil).PremiereDate).To(Equal("1999-01-01T00:00:00.0000000Z")) }) It("is omitted when the track has no date at all", func() { @@ -449,8 +452,8 @@ var _ = Describe("mappers", func() { }) It("is set on albums from their date, falling back to MaxYear", func() { - Expect(*AlbumToBaseItem(model.Album{ID: testID("a1"), Date: "2013-09-06"}, nil).PremiereDate).To(Equal("2013-09-06T00:00:00Z")) - Expect(*AlbumToBaseItem(model.Album{ID: testID("a2"), MaxYear: 2013}, nil).PremiereDate).To(Equal("2013-01-01T00:00:00Z")) + Expect(*AlbumToBaseItem(model.Album{ID: testID("a1"), Date: "2013-09-06"}, nil).PremiereDate).To(Equal("2013-09-06T00:00:00.0000000Z")) + Expect(*AlbumToBaseItem(model.Album{ID: testID("a2"), MaxYear: 2013}, nil).PremiereDate).To(Equal("2013-01-01T00:00:00.0000000Z")) Expect(AlbumToBaseItem(model.Album{ID: testID("a3")}, nil).PremiereDate).To(BeNil()) }) }) @@ -475,6 +478,13 @@ var _ = Describe("mappers", func() { Expect(item.ImageBlurHashes).To(BeNil()) }) + It("emits DateCreated and SortName on playlists, like albums and artists", func() { + p := model.Playlist{ID: testID("pl-1"), Name: "Chill", CreatedAt: time.Date(2026, 7, 1, 12, 0, 0, 0, time.UTC)} + Expect(PlaylistToBaseItem(p, nil).DateCreated).To(Equal("2026-07-01T12:00:00.0000000Z")) + Expect(PlaylistToBaseItem(p, nil).SortName).To(BeEmpty()) + Expect(PlaylistToBaseItem(p, ParseFields("SortName")).SortName).To(Equal("Chill")) + }) + It("changes the playlist image tag when the cover content changes", func() { p := model.Playlist{ID: testID("pl-1"), Name: "Chill"} p.ImageHash = "1111111111111111" 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/browsing_test.go b/server/jellyfin/e2e/browsing_test.go index 832b69f25..68d87cd25 100644 --- a/server/jellyfin/e2e/browsing_test.go +++ b/server/jellyfin/e2e/browsing_test.go @@ -1,6 +1,7 @@ package e2e import ( + "encoding/json" "net/http" "slices" "sort" @@ -110,10 +111,51 @@ var _ = Describe("Browsing", func() { Expect(q.Items).To(BeEmpty()) }) - It("defaults to albums when IncludeItemTypes is unrecognized", func() { + // Manet syncs collections as Boxset, and took albums coming back instead as a sync failure. + DescribeTable("returns nothing for a type it does not serve", + func(itemType string) { + q := queryResult(get("/Items?IncludeItemTypes=" + itemType + "&Recursive=true")) + Expect(q.TotalRecordCount).To(Equal(0)) + Expect(q.Items).To(BeEmpty()) + }, + Entry("a Jellyfin kind Navidrome has none of", "Boxset"), + ) + + It("treats a name Jellyfin does not know as an absent IncludeItemTypes", func() { q := queryResult(get("/Items?IncludeItemTypes=Nonsense&Recursive=true")) - Expect(q.TotalRecordCount).To(Equal(5)) + Expect(q.TotalRecordCount).To(Equal(queryResult(get("/Items?Recursive=true")).TotalRecordCount)) + Expect(q.TotalRecordCount).To(BeNumerically(">", 0)) }) + + // A strict client (Manet) fails its whole sync on the first item missing any of these. + DescribeTable("sends the keys Jellyfin puts on every item", + func(itemType string, fields string, nonNull ...string) { + var body struct{ Items []map[string]json.RawMessage } + res := get("/Items?IncludeItemTypes=" + itemType + "&Fields=" + fields + "&Recursive=true") + Expect(json.Unmarshal(res.Body.Bytes(), &body)).To(Succeed()) + Expect(body.Items).ToNot(BeEmpty()) + for _, it := range body.Items { + Expect(it).To(HaveKey("ChannelId"), "ChannelId is null, but always present") + for _, k := range nonNull { + Expect(it).To(HaveKey(k)) + Expect(string(it[k])).ToNot(Equal("null"), k) + } + } + }, + Entry("songs", "Audio", "Genres,Tags", "ImageTags", "HasLyrics", "Genres", "GenreItems", "Tags"), + Entry("albums", "MusicAlbum", "Genres", "ImageTags", "Genres", "GenreItems"), + Entry("artists", "MusicArtist", "Genres", "ImageTags", "Genres", "GenreItems"), + ) + + It("sends MediaType Unknown on items without one, as Jellyfin always emits it", func() { + q := queryResult(get("/Items?IncludeItemTypes=MusicAlbum&Recursive=true")) + Expect(q.Items).ToNot(BeEmpty()) + for _, it := range q.Items { + Expect(it.MediaType).To(Equal("Unknown")) + } + Expect(queryResult(get("/Items?IncludeItemTypes=Audio&Recursive=true")).Items[0].MediaType).To(Equal("Audio")) + }) + }) Describe("ParentId browsing", func() { diff --git a/server/jellyfin/e2e/discovery_test.go b/server/jellyfin/e2e/discovery_test.go new file mode 100644 index 000000000..0126ef3c6 --- /dev/null +++ b/server/jellyfin/e2e/discovery_test.go @@ -0,0 +1,50 @@ +package e2e + +import ( + "context" + "encoding/json" + "net" + "time" + + "github.com/navidrome/navidrome/server/jellyfin" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Auto discovery", func() { + BeforeEach(func() { setupTestDB() }) + + It("advertises the same Id and Name as /System/Info/Public", func() { + server, err := net.ListenPacket("udp4", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + client, err := net.ListenPacket("udp4", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(client.Close) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + jellyfin.NewDiscovery(ds).ServeOn(ctx, server) + }() + DeferCleanup(func() { + cancel() + Eventually(done).Should(BeClosed()) + }) + + _, err = client.WriteTo([]byte("who is JellyfinServer?"), server.LocalAddr()) + Expect(err).ToNot(HaveOccurred()) + Expect(client.SetReadDeadline(time.Now().Add(time.Second))).To(Succeed()) + buf := make([]byte, 1024) + n, _, err := client.ReadFrom(buf) + Expect(err).ToNot(HaveOccurred()) + var reply map[string]any + Expect(json.Unmarshal(buf[:n], &reply)).To(Succeed()) + + var pub map[string]any + parseInto(rawReq("GET", "/System/Info/Public", ""), &pub) + Expect(reply["Id"]).To(Equal(pub["Id"])) + Expect(reply["Name"]).To(Equal(pub["ServerName"])) + Expect(reply["Address"]).To(HaveSuffix("/jellyfin")) + }) +}) diff --git a/server/jellyfin/e2e/e2e_suite_test.go b/server/jellyfin/e2e/e2e_suite_test.go index 1ee19ea01..4f3cd82b5 100644 --- a/server/jellyfin/e2e/e2e_suite_test.go +++ b/server/jellyfin/e2e/e2e_suite_test.go @@ -44,6 +44,7 @@ import ( "github.com/navidrome/navidrome/core/lyrics" "github.com/navidrome/navidrome/core/matcher" "github.com/navidrome/navidrome/core/playlists" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/core/scrobbler" "github.com/navidrome/navidrome/core/sonic" "github.com/navidrome/navidrome/core/storage/storagetest" @@ -221,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()) @@ -237,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 { @@ -249,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 { @@ -261,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 { @@ -273,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 { @@ -338,6 +345,7 @@ func setupTestDB() { sonicSvc, lyrics.NewLyrics(ds, nil), events.NoopBroker(), + quickconnect.New(), ) } @@ -391,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/images_test.go b/server/jellyfin/e2e/images_test.go index 7fb85cba8..d1c3c132b 100644 --- a/server/jellyfin/e2e/images_test.go +++ b/server/jellyfin/e2e/images_test.go @@ -41,6 +41,10 @@ var _ = Describe("Item images", func() { Expect(u.IsAdmin).To(BeTrue()) }) + It("answers HEAD without authentication", func() { + Expect(rawReq("HEAD", "/Items/"+enc(albumID("Abbey Road"))+"/Images/Primary", "").Code).To(Equal(http.StatusOK)) + }) + It("serves images without authentication (public route)", func() { id := albumID("IV") w := rawReq("GET", "/Items/"+enc(id)+"/Images/Primary", "") diff --git a/server/jellyfin/e2e/lyrics_test.go b/server/jellyfin/e2e/lyrics_test.go index a49019bbf..d52a3b00f 100644 --- a/server/jellyfin/e2e/lyrics_test.go +++ b/server/jellyfin/e2e/lyrics_test.go @@ -58,14 +58,14 @@ var _ = Describe("Lyrics", func() { }) Describe("HasLyrics badge", func() { - It("is true for a track with embedded lyrics and omitted/false otherwise", func() { + It("is true for a track with embedded lyrics and false otherwise", func() { var stairway dto.BaseItemDto parseInto(get("/Items/"+enc(songID("Stairway To Heaven"))), &stairway) - Expect(stairway.HasLyrics).To(BeTrue()) + Expect(*stairway.HasLyrics).To(BeTrue()) var soWhat dto.BaseItemDto parseInto(get("/Items/"+enc(songID("So What"))), &soWhat) - Expect(soWhat.HasLyrics).To(BeFalse()) + Expect(*soWhat.HasLyrics).To(BeFalse()) }) }) }) 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 50c2e17f2..49c8e5a8c 100644 --- a/server/jellyfin/e2e/playlists_test.go +++ b/server/jellyfin/e2e/playlists_test.go @@ -19,13 +19,21 @@ var _ = Describe("Playlists", func() { playlistItems := func(plID string) dto.QueryResult { return queryResult(get("/Playlists/" + enc(plID) + "/Items")) } + 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()) }) @@ -47,16 +55,24 @@ 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() { - 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)) }) @@ -102,6 +118,36 @@ var _ = Describe("Playlists", func() { Expect(playlistItems(plID).TotalRecordCount).To(BeZero()) }) + Describe("at a position", func() { + var plID string + addAt := func(position string, ids ...string) { + url := "/Playlists/" + enc(plID) + "/Items?position=" + position + for _, id := range ids { + url += "&ids=" + enc(id) + } + Expect(post(url, "").Code).To(Equal(http.StatusNoContent)) + } + BeforeEach(func() { + plID = createPlaylist("Insert", []string{enc(songID("Come Together")), enc(songID("Something"))}) + }) + + DescribeTable("inserts at Jellyfin's zero-based position", + func(position string, want []string) { + addAt(position, songID("So What"), songID("Help!")) + Expect(order(plID)).To(Equal(want)) + }, + Entry("in the middle", "1", []string{"Come Together", "So What", "Help!", "Something"}), + Entry("first, for zero or less", "-3", []string{"So What", "Help!", "Come Together", "Something"}), + Entry("last, past the end", "9", []string{"Come Together", "Something", "So What", "Help!"}), + ) + + It("keeps an expanded album's track order", func() { + albumOrder := order(createPlaylist("Album", []string{enc(albumID("Abbey Road"))})) + addAt("1", albumID("Abbey Road")) + Expect(order(plID)).To(Equal(append(append([]string{"Come Together"}, albumOrder...), "Something"))) + }) + }) + It("removes an entry by its PlaylistItemId", func() { plID := createPlaylist("Remove", []string{enc(songID("Come Together")), enc(songID("Something"))}) entryID := playlistItems(plID).Items[0].PlaylistItemId @@ -118,6 +164,58 @@ var _ = Describe("Playlists", func() { }) }) + Describe("move", func() { + move := func(u model.User, plID, entryID string, newIndex string) int { + return jReq(u, "POST", "/Playlists/"+enc(plID)+"/Items/"+entryID+"/Move/"+newIndex, "").Code + } + var plID string + var entries []dto.BaseItemDto + BeforeEach(func() { + plID = createPlaylist("Move", []string{enc(songID("Come Together")), enc(songID("Something")), enc(songID("So What"))}) + entries = playlistItems(plID).Items + }) + + DescribeTable("moves an entry to Jellyfin's zero-based index", + func(entry int, newIndex string, want []string) { + Expect(move(adminUser, plID, entries[entry].PlaylistItemId, newIndex)).To(Equal(http.StatusNoContent)) + Expect(order(plID)).To(Equal(want)) + var positions []string + for _, it := range playlistItems(plID).Items { + positions = append(positions, it.PlaylistItemId) + } + Expect(positions).To(Equal([]string{dto.EncodePlaylistEntryID("1"), dto.EncodePlaylistEntryID("2"), dto.EncodePlaylistEntryID("3")})) + }, + Entry("towards the end", 0, "2", []string{"Something", "So What", "Come Together"}), + Entry("towards the start", 2, "0", []string{"So What", "Come Together", "Something"}), + Entry("to the end, past the last index", 0, "99", []string{"Something", "So What", "Come Together"}), + Entry("to the end, for the largest int", 0, "9223372036854775807", []string{"Something", "So What", "Come Together"}), + ) + + It("ignores an entry that is not in the playlist", func() { + Expect(move(adminUser, plID, dto.EncodePlaylistEntryID("42"), "0")).To(Equal(http.StatusNoContent)) + Expect(order(plID)).To(Equal([]string{"Come Together", "Something", "So What"})) + }) + + It("404s on a malformed entry id", func() { + Expect(move(adminUser, plID, enc(songID("So What")), "0")).To(Equal(http.StatusNotFound)) + }) + + It("hides another user's private playlist", func() { + Expect(move(regularUser, plID, entries[0].PlaylistItemId, "2")).To(Equal(http.StatusNotFound)) + Expect(order(plID)).To(Equal([]string{"Come Together", "Something", "So What"})) + }) + + It("forbids a non-owner on a public playlist", func() { + Expect(post("/Playlists/"+enc(plID), `{"Name":"Move","IsPublic":true}`).Code).To(Equal(http.StatusNoContent)) + Expect(move(regularUser, plID, entries[0].PlaylistItemId, "2")).To(Equal(http.StatusForbidden)) + Expect(order(plID)).To(Equal([]string{"Come Together", "Something", "So What"})) + }) + + It("rejects a negative index", func() { + Expect(move(adminUser, plID, entries[0].PlaylistItemId, "-1")).To(Equal(http.StatusBadRequest)) + }) + }) + Describe("users", func() { It("reports the current user as an editor", func() { plID := createPlaylist("Perms", nil) @@ -220,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()) }) @@ -240,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()) @@ -272,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)) }) @@ -283,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")) }) @@ -315,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()) }) @@ -329,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/quickconnect_test.go b/server/jellyfin/e2e/quickconnect_test.go new file mode 100644 index 000000000..1beffc56f --- /dev/null +++ b/server/jellyfin/e2e/quickconnect_test.go @@ -0,0 +1,91 @@ +package e2e + +import ( + "net/http" + "net/http/httptest" + "strings" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/server/jellyfin/dto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("QuickConnect", func() { + BeforeEach(func() { + setupTestDB() + DeferCleanup(configtest.SetupConfig()) + conf.Server.Jellyfin.QuickConnect = true + }) + + clientReq := func(deviceID, method, path, body string) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + r := httptest.NewRequest(method, path, strings.NewReader(body)) + r.Header.Set("Authorization", `MediaBrowser Client="Finamp", Device="Pixel 7", DeviceId="`+deviceID+`", Version="1.0"`) + r.Header.Set("Content-Type", "application/json") + router.ServeHTTP(w, r) + return w + } + initiate := func() dto.QuickConnectResult { + var pending dto.QuickConnectResult + parseInto(clientReq("new-device", "POST", "/QuickConnect/Initiate", ""), &pending) + Expect(pending.Authenticated).To(BeFalse()) + return pending + } + redeem := func(deviceID, secret string) *httptest.ResponseRecorder { + return clientReq(deviceID, "POST", "/Users/AuthenticateWithQuickConnect", `{"Secret":"`+secret+`"}`) + } + + It("signs a new client in with a code approved by a signed-in user", func() { + Expect(rawReq("GET", "/QuickConnect/Enabled", "").Body.String()).To(MatchJSON("true")) + pending := initiate() + + Expect(postAs(regularUser, "/QuickConnect/Authorize?Code="+pending.Code, "").Code).To(Equal(http.StatusOK)) + + var status dto.QuickConnectResult + parseInto(rawReq("GET", "/QuickConnect/Connect?Secret="+pending.Secret, ""), &status) + Expect(status.Authenticated).To(BeTrue()) + + // Android TV redeems with a different DeviceId than it initiated with. + w := redeem("other-device", pending.Secret) + Expect(w.Code).To(Equal(http.StatusOK)) + var res dto.AuthenticationResult + parseInto(w, &res) + Expect(res.User.Name).To(Equal(regularUser.UserName)) + Expect(res.SessionInfo).ToNot(BeNil()) + Expect(res.SessionInfo.DeviceId).To(Equal("other-device")) + + r := httptest.NewRequest("GET", "/Users/Me", nil) + r.Header.Set("X-Emby-Token", res.AccessToken) + me := httptest.NewRecorder() + router.ServeHTTP(me, r) + Expect(me.Code).To(Equal(http.StatusOK)) + + Expect(redeem("new-device", pending.Secret).Code).To(Equal(http.StatusNotFound)) + }) + + It("lets an admin approve a code for another user", func() { + pending := initiate() + Expect(post("/QuickConnect/Authorize?Code="+pending.Code+"&UserId="+enc(regularUser.ID), "").Code). + To(Equal(http.StatusOK)) + + var res dto.AuthenticationResult + parseInto(redeem("new-device", pending.Secret), &res) + Expect(res.User.Name).To(Equal(regularUser.UserName)) + }) + + It("requires authentication to approve a code", func() { + pending := initiate() + Expect(rawReq("POST", "/QuickConnect/Authorize?Code="+pending.Code, "").Code).To(Equal(http.StatusUnauthorized)) + }) + + It("answers 401 on every Quick Connect call when disabled", func() { + conf.Server.Jellyfin.QuickConnect = false + Expect(rawReq("GET", "/QuickConnect/Enabled", "").Body.String()).To(MatchJSON("false")) + Expect(clientReq("new-device", "POST", "/QuickConnect/Initiate", "").Code).To(Equal(http.StatusUnauthorized)) + Expect(rawReq("GET", "/QuickConnect/Connect?Secret=x", "").Code).To(Equal(http.StatusUnauthorized)) + Expect(post("/QuickConnect/Authorize?Code=123456", "").Code).To(Equal(http.StatusUnauthorized)) + Expect(redeem("new-device", "x").Code).To(Equal(http.StatusUnauthorized)) + }) +}) diff --git a/server/jellyfin/e2e/sessions_test.go b/server/jellyfin/e2e/sessions_test.go index 2057b43c6..22c87b108 100644 --- a/server/jellyfin/e2e/sessions_test.go +++ b/server/jellyfin/e2e/sessions_test.go @@ -27,16 +27,20 @@ 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)) }) + It("acknowledges a playback session ping", func() { + Expect(post("/Sessions/Playing/Ping?playSessionId=abc", "").Code).To(Equal(http.StatusNoContent)) + }) + It("does not count a brief play stopped before the threshold", func() { // Regression: Finamp sends a Stopped report on every track switch, so an immediate skip // (1 second in) must not mark the track played. Seeded tracks are >= 120s, so the 50% @@ -44,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 5f1a88a21..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 @@ -150,4 +150,33 @@ var _ = Describe("Similar", func() { Expect(w.Code).To(Equal(http.StatusNotFound)) }) }) + + Describe("type-specific InstantMix routes", func() { + BeforeEach(func() { + providerFake.similarSongs = model.MediaFiles{{ID: testID("x1"), Title: "Mix Song", LibraryID: 1}} + }) + + It("leads a song mix with the seed on /Songs/{id}/InstantMix", func() { + q := queryResult(get("/Songs/" + enc(songID("So What")) + "/InstantMix")) + Expect(names(q.Items)).To(Equal([]string{"So What", "Mix Song"})) + }) + + DescribeTable("returns the provider's mix", + func(path func() string) { + q := queryResult(get(path())) + Expect(names(q.Items)).To(Equal([]string{"Mix Song"})) + }, + Entry("Albums/{id}", func() string { return "/Albums/" + enc(albumID("Abbey Road")) + "/InstantMix" }), + Entry("Artists/{id}", func() string { return "/Artists/" + enc(artistID("Miles Davis")) + "/InstantMix" }), + Entry("Playlists/{id}", func() string { + return "/Playlists/" + enc(createPlaylist("Seed", []string{enc(songID("So What"))})) + "/InstantMix" + }), + Entry("Artists/InstantMix?id=", func() string { return "/Artists/InstantMix?id=" + enc(artistID("Miles Davis")) }), + Entry("MusicGenres/InstantMix?id=", func() string { return "/MusicGenres/InstantMix?id=" + enc(genreID("Jazz")) }), + ) + + It("404s a malformed id query param", func() { + Expect(get("/Artists/InstantMix?id=not-a-valid-id").Code).To(Equal(http.StatusNotFound)) + }) + }) }) diff --git a/server/jellyfin/e2e/streaming_test.go b/server/jellyfin/e2e/streaming_test.go index 67208aed3..9b7380185 100644 --- a/server/jellyfin/e2e/streaming_test.go +++ b/server/jellyfin/e2e/streaming_test.go @@ -1,6 +1,7 @@ package e2e import ( + "fmt" "net/http" "strings" @@ -45,6 +46,19 @@ var _ = Describe("Streaming", func() { It("returns 404 for an unknown track", func() { Expect(get("/Audio/" + enc(testID("nope")) + "/stream").Code).To(Equal(http.StatusNotFound)) }) + + DescribeTable("answers HEAD (Fintunes probes the type before playing or downloading)", + func(path string) { + w := jReq(adminUser, "HEAD", fmt.Sprintf(path, enc(songID("Help!"))), "") + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).ToNot(BeEmpty()) + }, + Entry("stream", "/Audio/%s/stream"), + Entry("stream.{container}", "/Audio/%s/stream.mp3"), + Entry("universal", "/Audio/%s/universal"), + Entry("File", "/Items/%s/File"), + Entry("Download", "/Items/%s/Download"), + ) }) Describe("GET /Audio/{id}/main.m3u8 (Finamp transcoding mode)", func() { diff --git a/server/jellyfin/e2e/system_test.go b/server/jellyfin/e2e/system_test.go index ac8c350c4..51443bdc9 100644 --- a/server/jellyfin/e2e/system_test.go +++ b/server/jellyfin/e2e/system_test.go @@ -81,12 +81,4 @@ var _ = Describe("System", func() { Expect(strings.TrimSpace(w.Body.String())).To(HavePrefix("Navidrome")) }) }) - - Describe("GET /QuickConnect/Enabled", func() { - It("reports QuickConnect disabled", func() { - w := rawReq("GET", "/QuickConnect/Enabled", "") - Expect(w.Code).To(Equal(http.StatusOK)) - Expect(strings.TrimSpace(w.Body.String())).To(Equal("false")) - }) - }) }) diff --git a/server/jellyfin/images.go b/server/jellyfin/images.go index 9f1b9ce1a..6af533dbb 100644 --- a/server/jellyfin/images.go +++ b/server/jellyfin/images.go @@ -23,14 +23,17 @@ import ( _ "golang.org/x/image/webp" ) -// imageSize picks the tighter of Jellyfin's two bounds, because Navidrome resizes on a single -// dimension: reading only MaxWidth serves the full-size original to a client that sent MaxHeight. -func imageSize(maxWidth, maxHeight int) int { - w, h := max(maxWidth, 0), max(maxHeight, 0) - if w == 0 || h == 0 { - return max(w, h) +// imageSize reduces Jellyfin's size params to Navidrome's single bound. Width/Height/Max* are +// bounds, so the tightest wins; Fill* must cover its box, so its larger side is the bound. +func imageSize(p *req.Values) int { + fill := max(p.IntOr("fillwidth", 0), p.IntOr("fillheight", 0)) + size := 0 + for _, v := range []int{p.IntOr("width", 0), p.IntOr("height", 0), p.IntOr("maxwidth", 0), p.IntOr("maxheight", 0), fill} { + if v > 0 && (size == 0 || v < size) { + size = v + } } - return min(w, h) + return size } func (api *Router) getItemImage(w http.ResponseWriter, r *http.Request) { @@ -41,8 +44,7 @@ func (api *Router) getItemImage(w http.ResponseWriter, r *http.Request) { if !ok { return } - p := req.Params(r) - size := imageSize(p.IntOr("maxwidth", 0), p.IntOr("maxheight", 0)) + size := imageSize(req.Params(r)) artID := api.resolveArtworkID(ctx, itemId) img, err := api.artwork.GetOrPlaceholder(ctx, artID, size, false) @@ -80,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 2f32e6ec9..b7cf9b7e7 100644 --- a/server/jellyfin/images_test.go +++ b/server/jellyfin/images_test.go @@ -65,10 +65,10 @@ func newImageRequest(itemId string) (*httptest.ResponseRecorder, *http.Request) var _ = Describe("Images", func() { // Real Jellyfin fits the image inside either bound, so a client that sends only MaxHeight must // still get a resized image rather than the full-size original. - DescribeTable("derives the requested size from MaxWidth or MaxHeight", + 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} @@ -84,11 +84,20 @@ var _ = Describe("Images", func() { Entry("both, smaller bound wins", "maxwidth=200&maxheight=300", 200), Entry("both, smaller bound wins regardless of order", "maxwidth=300&maxheight=200", 200), Entry("neither", "", 0), + Entry("Width only", "width=300", 300), + Entry("Height only", "height=300", 300), + Entry("Width capped by MaxWidth", "width=500&maxwidth=300", 300), + Entry("FillWidth only", "fillwidth=300", 300), + Entry("FillHeight only", "fillheight=300", 300), + Entry("both fill sides, larger one covers the box", "fillwidth=200&fillheight=300", 300), + Entry("fill smaller than MaxWidth wins", "maxwidth=500&fillwidth=300&fillheight=300", 300), + Entry("MaxWidth smaller than fill wins", "maxwidth=200&fillwidth=300&fillheight=300", 200), + Entry("non-positive values are ignored", "fillwidth=0&maxwidth=-5&height=250", 250), ) 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} @@ -129,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} @@ -144,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} @@ -160,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} @@ -176,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} @@ -194,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 d16c5e57e..28de4c31e 100644 --- a/server/jellyfin/items.go +++ b/server/jellyfin/items.go @@ -219,11 +219,40 @@ func (api *Router) writeItemsArray(w http.ResponseWriter, r *http.Request, res i api.streamResult(w, r, res, streamItemsArray) } -// streamResult stamps every item's ServerId (constant per request, so it's set here rather than in -// each mapper). The cursor opens before the first byte, so a failed open is still a clean 500. +// stampItem fills in what real Jellyfin puts on every item it returns, so a client that requires a +// key never meets an item without it: ServerId, MediaType, ImageTags, and the Fields-gated lists. +func stampItem(it dto.BaseItemDto, serverID string, fields dto.Fields) dto.BaseItemDto { + it.ServerId = serverID + if it.MediaType == "" { + it.MediaType = "Unknown" + } + if it.ImageTags == nil { + it.ImageTags = map[string]string{} + } + if fields.Has("Genres") { + if it.Genres == nil { + it.Genres = []string{} + } + if it.GenreItems == nil { + it.GenreItems = []dto.NameGuidPair{} + } + } + if fields.Has("Tags") && it.Tags == nil { + it.Tags = []string{} + } + return it +} + +// requestFields parses the Fields param, which gates what stampItem and the mappers attach. +func requestFields(r *http.Request) dto.Fields { + return dto.ParseFields(req.Params(r).Strings("fields")...) +} + +// streamResult stamps every item (see stampItem). The cursor opens before the first byte, so a +// failed open is still a clean 500. func (api *Router) streamResult(w http.ResponseWriter, r *http.Request, res itemsResult, write func(io.Writer, iter.Seq2[dto.BaseItemDto, error]) error) { - sid := api.serverID(r.Context()) + sid, fields := api.serverID(r.Context()), requestFields(r) seq, err := res.seq() if err != nil { api.internalError(w, r, err) @@ -235,8 +264,7 @@ func (api *Router) streamResult(w http.ResponseWriter, r *http.Request, res item yield(dto.BaseItemDto{}, err) return } - it.ServerId = sid - if !yield(it, nil) { + if !yield(stampItem(it, sid, fields), nil) { return } } @@ -321,7 +349,7 @@ func (api *Router) parseItemsQuery(ctx context.Context, r *http.Request) (itemsQ } q := listParams(p) q.ids = ids - q.rawTypes = p.StringOr("includeitemtypes", "") + q.rawTypes = knownItemKinds(p.StringOr("includeitemtypes", "")) q.parentId = parentId q.genreIds = genreIds q.albumIds = albumIds @@ -354,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"} } } @@ -380,11 +408,11 @@ func (api *Router) queryItems(ctx context.Context, r *http.Request) (itemsResult case len(q.ids) > 0: return materialized(api.itemsByIDs(ctx, q.ids, q.fields)), nil // A ManualPlaylistsFolder query asks for the synthetic "playlists library" container, not real items. - case strings.Contains(q.rawTypes, "ManualPlaylistsFolder"): + case strings.Contains(strings.ToLower(q.rawTypes), "manualplaylistsfolder"): 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) @@ -559,23 +587,49 @@ func parseYears(r *http.Request) []int { return years } -// parseTypes returns the recognized entries in IncludeItemTypes in order, defaulting to -// {"MusicAlbum"} when none are recognized (so ParentId= browses that artist's albums). +// supportedTypes maps lowercased IncludeItemTypes names (Jellyfin binds them case-insensitively) +// to the item types Navidrome serves. +var supportedTypes = map[string]string{ + "audio": "Audio", "musicartist": "MusicArtist", "musicalbum": "MusicAlbum", "musicgenre": "MusicGenre", "playlist": "Playlist", +} + +// jellyfinItemKinds lists Jellyfin's BaseItemKind names, lowercased. +var jellyfinItemKinds = map[string]bool{ + "aggregatefolder": true, "audio": true, "audiobook": true, "basepluginfolder": true, "book": true, + "boxset": true, "channel": true, "channelfolderitem": true, "collectionfolder": true, "episode": true, + "folder": true, "genre": true, "manualplaylistsfolder": true, "movie": true, "livetvchannel": true, + "livetvprogram": true, "musicalbum": true, "musicartist": true, "musicgenre": true, "musicvideo": true, + "person": true, "photo": true, "photoalbum": true, "playlist": true, "playlistsfolder": true, + "program": true, "recording": true, "season": true, "series": true, "studio": true, "trailer": true, + "tvchannel": true, "tvprogram": true, "userrootfolder": true, "userview": true, "video": true, "year": true, +} + +// knownItemKinds drops IncludeItemTypes entries that aren't BaseItemKind names, as Jellyfin's binder +// does, so an all-unknown list (JellyBox sends "music") behaves like an absent one. +func knownItemKinds(types string) string { + var known []string + for t := range strings.SplitSeq(types, ",") { + if t = strings.TrimSpace(t); jellyfinItemKinds[strings.ToLower(t)] { + known = append(known, t) + } + } + return strings.Join(known, ",") +} + +// parseTypes returns the supported entries in IncludeItemTypes in order. Only an absent param +// defaults to albums, so ParentId= still browses that artist's albums. func parseTypes(types string) []string { + if strings.TrimSpace(types) == "" { + return []string{"MusicAlbum"} + } var recognized []string for t := range strings.SplitSeq(types, ",") { - t = strings.TrimSpace(t) - switch t { - case "Audio", "MusicArtist", "MusicAlbum", "MusicGenre", "Playlist": - recognized = append(recognized, t) + if name, ok := supportedTypes[strings.ToLower(strings.TrimSpace(t))]; ok { + recognized = append(recognized, name) } } // Dedupe: a repeated type would duplicate items in the merge and spawn a redundant query. - recognized = slice.Unique(recognized) - if len(recognized) == 0 { - return []string{"MusicAlbum"} - } - return recognized + return slice.Unique(recognized) } // paginate applies StartIndex/Limit to an in-memory item list, for the multi-type merge path only @@ -648,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). @@ -679,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 { @@ -728,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 @@ -741,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 } @@ -752,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 @@ -764,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 @@ -780,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 } @@ -791,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 } @@ -805,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 } @@ -828,28 +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) { - for _, lib := range u.Libraries { - if lib.ID == libID { - return libraryView(lib), true - } - } - // Admin bypass: Libraries is empty but all access is granted, so fetch the real library. - if lib, err := api.ds.Library(ctx).Get(libID); err == nil { - return libraryView(*lib), true + 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 { - // 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. + 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 } @@ -859,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 @@ -870,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 @@ -950,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 17180c36a..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) @@ -74,6 +74,18 @@ var _ = Describe("Items", func() { Expect(res.Items[0].Id).To(Equal(dto.EncodeID(testID("s1")))) }) + It("ignores IncludeItemTypes names that aren't Jellyfin item kinds", func() { + 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) + var res dto.QueryResult + Expect(json.Unmarshal(w.Body.Bytes(), &res)).To(Succeed()) + Expect(res.Items).To(HaveLen(1)) + Expect(res.Items[0].Id).To(Equal(dto.EncodeID(testID("s1")))) + }) + It("lists a playlist's tracks when ParentId is a playlist, whatever the type", func() { fp.getPls = &model.Playlist{ID: testID("pl1"), Tracks: model.PlaylistTracks{ {ID: "1", MediaFileID: testID("s1"), PlaylistID: testID("pl1"), MediaFile: model.MediaFile{ID: testID("s1")}}, @@ -114,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()) @@ -128,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) @@ -139,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() { @@ -211,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) @@ -246,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) @@ -260,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) @@ -275,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()) @@ -300,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) @@ -314,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()) @@ -329,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()) @@ -369,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()) @@ -378,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 @@ -398,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()) @@ -410,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()) @@ -420,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). @@ -431,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). @@ -442,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). @@ -460,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()) @@ -474,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()) @@ -493,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), @@ -514,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), @@ -534,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()) @@ -553,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), @@ -569,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() @@ -583,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()) @@ -597,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) @@ -611,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) @@ -628,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()) @@ -644,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) @@ -661,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()) @@ -698,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", @@ -716,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}} @@ -730,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}} @@ -744,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}} @@ -758,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}} @@ -773,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 @@ -791,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()) @@ -812,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) @@ -820,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) @@ -828,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) @@ -836,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()) @@ -848,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()) @@ -863,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() { @@ -915,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) @@ -926,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"))) @@ -947,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"))) @@ -956,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"))) @@ -965,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"))) @@ -981,6 +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().(*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) @@ -1030,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"))) @@ -1044,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)) @@ -1060,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) @@ -1072,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}} @@ -1124,6 +1137,25 @@ var _ = Describe("Items", func() { It("dedupes repeated types, preserving first-seen order", func() { Expect(parseTypes("Audio,MusicAlbum,Audio")).To(Equal([]string{"Audio", "MusicAlbum"})) }) + + It("matches type names case-insensitively, like Jellyfin's enum binding", func() { + Expect(parseTypes("musicalbum, AUDIO")).To(Equal([]string{"MusicAlbum", "Audio"})) + }) + + It("returns no types for real Jellyfin kinds Navidrome has none of", func() { + Expect(parseTypes("Boxset")).To(BeEmpty()) + Expect(parseTypes("BoxSet,Movie")).To(BeEmpty()) + }) + + It("drops names that aren't Jellyfin item kinds, case-insensitively", func() { + Expect(knownItemKinds("music")).To(BeEmpty()) + Expect(knownItemKinds("music, audio,MUSICVIDEO")).To(Equal("audio,MUSICVIDEO")) + }) + + It("defaults to albums only when IncludeItemTypes is absent", func() { + Expect(parseTypes("")).To(Equal([]string{"MusicAlbum"})) + Expect(parseTypes("Nonsense")).To(BeEmpty()) + }) }) Describe("decodeFilterParam", func() { diff --git a/server/jellyfin/library.go b/server/jellyfin/library.go index b21603843..9f1f7ee73 100644 --- a/server/jellyfin/library.go +++ b/server/jellyfin/library.go @@ -7,7 +7,6 @@ import ( "github.com/Masterminds/squirrel" "github.com/go-chi/chi/v5" - "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" "github.com/navidrome/navidrome/server/jellyfin/dto" "github.com/navidrome/navidrome/utils/req" @@ -74,16 +73,3 @@ func libraryScopeFilter(scope []int) squirrel.Sqlizer { } return squirrel.Eq{"library_tag.library_id": scope} } - -// libraryView builds the CollectionFolder BaseItemDto representing a library as a top-level node. -// Shared by getUserViews and getItem, since Finamp fetches a UserView's id as a plain item. -func libraryView(lib model.Library) dto.BaseItemDto { - return dto.BaseItemDto{ - Id: dto.EncodeLibraryID(lib.ID), - Name: lib.Name, - Type: "CollectionFolder", - CollectionType: "music", - IsFolder: true, - BackdropImageTags: []string{}, - } -} 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 ac790a821..5569f4f83 100644 --- a/server/jellyfin/playlists.go +++ b/server/jellyfin/playlists.go @@ -4,7 +4,9 @@ import ( "context" "encoding/json" "errors" + "math" "net/http" + "strconv" "strings" "github.com/go-chi/chi/v5" @@ -46,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 @@ -67,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)}) } @@ -136,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 } @@ -178,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 @@ -206,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 @@ -239,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 })...) @@ -253,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 @@ -261,8 +271,8 @@ func (api *Router) songIDs(ctx context.Context, opts model.QueryOptions) []strin return slice.Map(mfs, func(mf model.MediaFile) string { return mf.ID }) } -// addToPlaylist appends items by id, expanding containers into tracks (see expandContainerIDs). -// AddTracks enforces ownership; a locked playlist maps to 403, any other error to 404. +// addToPlaylist adds items by id (containers expand to tracks), inserting at the zero-based position +// when given. Core enforces ownership: a locked playlist maps to 403, any other error to 404. func (api *Router) addToPlaylist(w http.ResponseWriter, r *http.Request) { ctx := r.Context() id, ok := itemIDParam(w, r, "playlistId") @@ -275,7 +285,13 @@ func (api *Router) addToPlaylist(w http.ResponseWriter, r *http.Request) { return } ids := api.expandContainerIDs(ctx, decoded) - if _, err := api.playlists.AddTracks(ctx, id, ids); err != nil { + var err error + if position, perr := req.Params(r).Int64("position"); perr == nil { + _, err = api.playlists.InsertTracks(ctx, id, ids, insertPosition(position)) + } else { + _, err = api.playlists.AddTracks(ctx, id, ids) + } + if err != nil { if errors.Is(err, model.ErrPlaylistNotEditable) { http.Error(w, "Forbidden", http.StatusForbidden) return @@ -286,6 +302,12 @@ func (api *Router) addToPlaylist(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) } +// insertPosition maps Jellyfin's zero-based position to a 1-based one, clamped in int64 first so +// it can't wrap on 32-bit builds. +func insertPosition(position int64) int { + return int(min(max(position, 0), math.MaxInt32-1) + 1) +} + // removeFromPlaylist removes entries by entryIds — playlist-entry ids (PlaylistItemId), not media // file ids, since RemoveTracks deletes playlist_tracks rows by that id. RemoveTracks enforces // ownership; a locked playlist maps to 403, any other error to 404. @@ -316,6 +338,38 @@ func (api *Router) removeFromPlaylist(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) } +// movePlaylistItem moves an entry to Jellyfin's zero-based newIndex. Reorder clamps past-the-end +// indexes and rejects unknown entries, which Jellyfin treats as a no-op. +func (api *Router) movePlaylistItem(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + id, ok := itemIDParam(w, r, "playlistId") + if !ok { + return + } + entry, ok := dto.DecodePlaylistEntryID(chi.URLParam(r, "entryId")) + if !ok { + http.Error(w, "Not Found", http.StatusNotFound) + return + } + newIndex, err := strconv.Atoi(chi.URLParam(r, "newIndex")) + if err != nil || newIndex < 0 { + http.Error(w, "Bad Request", http.StatusBadRequest) + return + } + // Resolve first, so a missing or hidden playlist is still a 404 below the unknown-entry no-op. + if _, err := api.playlists.Get(ctx, id); err != nil { + api.playlistError(w, r, err) + return + } + pos, _ := strconv.Atoi(entry) + err = api.playlists.ReorderTrack(ctx, id, pos, min(newIndex, math.MaxInt32-1)+1) + if err != nil && !errors.Is(err, model.ErrNotFound) { + api.playlistError(w, r, err) + return + } + w.WriteHeader(http.StatusNoContent) +} + // Clients probe these before offering edits. Navidrome has no per-playlist ACL, so CanEdit carries // only editability; ownership is enforced on write, and a lookup error 404s to prevent probing. func (api *Router) getPlaylistUsers(w http.ResponseWriter, r *http.Request) { diff --git a/server/jellyfin/playlists_test.go b/server/jellyfin/playlists_test.go index 270f5fe08..d38ef3b44 100644 --- a/server/jellyfin/playlists_test.go +++ b/server/jellyfin/playlists_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "io" + "math" "net/http" "net/http/httptest" "strings" @@ -54,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 { @@ -133,6 +144,17 @@ func (f *fakePlaylists) RemoveImage(_ context.Context, playlistID string) error return f.removeImageErr } +var _ = DescribeTable("insertPosition", + func(position int64, want int) { + Expect(insertPosition(position)).To(Equal(want)) + }, + Entry("zero-based index becomes a 1-based position", int64(2), 3), + Entry("negative prepends", int64(-5), 1), + Entry("beyond int32 stays past the end instead of wrapping", int64(1)<<32+1, math.MaxInt32), + Entry("largest int64 stays past the end", int64(math.MaxInt64), math.MaxInt32), + Entry("smallest int64 prepends", int64(math.MinInt64), 1), +) + var _ = Describe("Playlists", func() { var api *Router var fp *fakePlaylists @@ -172,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() { @@ -252,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 new file mode 100644 index 000000000..f7fa1b979 --- /dev/null +++ b/server/jellyfin/quickconnect.go @@ -0,0 +1,145 @@ +package jellyfin + +import ( + "encoding/json" + "errors" + "net/http" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/core/quickconnect" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/server/jellyfin/dto" +) + +func requireQuickConnect(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !conf.Server.Jellyfin.QuickConnect { + http.Error(w, "Quick connect is disabled", http.StatusUnauthorized) + return + } + next.ServeHTTP(w, r) + }) +} + +func (api *Router) quickConnectEnabled(w http.ResponseWriter, r *http.Request) { + api.ok(w, r, conf.Server.Jellyfin.QuickConnect) +} + +// Initiate is unauthenticated and its fields are kept for minutes, so their size must be bounded. +const maxQuickConnectField = 512 + +func (api *Router) quickConnectInitiate(w http.ResponseWriter, r *http.Request) { + a := parseMediaBrowserAuth(r) + if a.Client == "" || a.Device == "" || a.DeviceId == "" || a.Version == "" || + max(len(a.Client), len(a.Device), len(a.DeviceId), len(a.Version)) > maxQuickConnectField { + http.Error(w, "Client, Device, DeviceId and Version are required", http.StatusBadRequest) + return + } + req, err := api.quickConnect.Initiate(quickconnect.Device{ + ID: a.DeviceId, Name: a.Device, App: a.Client, AppVersion: a.Version, + }) + if errors.Is(err, quickconnect.ErrTooManyRequests) { + http.Error(w, "Too Many Requests", http.StatusTooManyRequests) + return + } + if err != nil { + api.internalError(w, r, err) + return + } + api.ok(w, r, quickConnectResult(req)) +} + +func (api *Router) quickConnectConnect(w http.ResponseWriter, r *http.Request) { + req, err := api.quickConnect.Status(r.URL.Query().Get("secret")) + if err != nil { + http.Error(w, "Unknown secret", http.StatusNotFound) + return + } + api.ok(w, r, quickConnectResult(req)) +} + +func (api *Router) quickConnectAuthorize(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + caller, _ := request.UserFrom(ctx) + userID := caller.ID + if encoded := r.URL.Query().Get("userid"); encoded != "" { + var ok bool + if userID, ok = dto.DecodeID(encoded); !ok { + http.Error(w, "Invalid userId", http.StatusBadRequest) + return + } + } + target := caller + if userID != caller.ID { + if !caller.IsAdmin { + http.Error(w, "Forbidden", http.StatusForbidden) + return + } + usr, err := api.ds.User().Get(ctx, userID) + if errors.Is(err, model.ErrNotFound) { + http.Error(w, "Unknown user", http.StatusNotFound) + return + } + if err != nil { + api.internalError(w, r, err) + return + } + target = *usr + } + + req, err := api.quickConnect.Authorize(r.URL.Query().Get("code"), target.ID) + switch { + case errors.Is(err, model.ErrNotFound): + http.Error(w, "Unknown code", http.StatusNotFound) + case errors.Is(err, quickconnect.ErrAlreadyAuthorized): + http.Error(w, "Code already used", http.StatusConflict) + case err != nil: + api.internalError(w, r, err) + default: + log.Info(ctx, "Jellyfin API: Quick Connect sign-in approved", "username", target.UserName, + "approvedBy", caller.UserName, "client", req.Device.App, "device", req.Device.Name) + api.ok(w, r, true) + } +} + +func (api *Router) authenticateWithQuickConnect(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + var body struct { + Secret string `json:"Secret"` + } + if err := json.NewDecoder(r.Body).Decode(&body); err != nil || body.Secret == "" { + http.Error(w, "Bad Request", http.StatusBadRequest) + return + } + userID, err := api.quickConnect.Redeem(body.Secret) + if err != nil { + http.Error(w, "Unknown secret", http.StatusNotFound) + return + } + 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) + return + } + if err != nil { + api.internalError(w, r, err) + return + } + api.signIn(w, r, usr) +} + +func quickConnectResult(req quickconnect.Request) dto.QuickConnectResult { + return dto.QuickConnectResult{ + Authenticated: req.Authorized(), + Secret: req.Secret, + Code: req.Code, + DeviceId: req.Device.ID, + DeviceName: req.Device.Name, + AppName: req.Device.App, + AppVersion: req.Device.AppVersion, + DateAdded: dto.JellyfinDate(&req.DateAdded), + } +} diff --git a/server/jellyfin/quickconnect_test.go b/server/jellyfin/quickconnect_test.go new file mode 100644 index 000000000..94573e6ff --- /dev/null +++ b/server/jellyfin/quickconnect_test.go @@ -0,0 +1,300 @@ +package jellyfin + +import ( + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/core/auth" + "github.com/navidrome/navidrome/core/quickconnect" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/server/jellyfin/dto" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type fullQuickConnect struct{ quickconnect.QuickConnect } + +func (fullQuickConnect) Initiate(quickconnect.Device) (quickconnect.Request, error) { + return quickconnect.Request{}, quickconnect.ErrTooManyRequests +} + +var _ = Describe("QuickConnect", func() { + const finamp = `MediaBrowser Client="Finamp", Device="Pixel 7", DeviceId="dev-1", Version="1.0.0"` + var ( + api *Router + ds *tests.MockDataStore + qc quickconnect.QuickConnect + alice = model.User{ID: testID("alice"), UserName: "alice"} + bob = model.User{ID: testID("bob"), UserName: "bob"} + admin = model.User{ID: testID("admin"), UserName: "admin", IsAdmin: true} + ) + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.Jellyfin.QuickConnect = true + ds = &tests.MockDataStore{} + auth.Init(ds) + ur := ds.User().(*tests.MockedUserRepo) + for _, u := range []model.User{alice, bob, admin} { + Expect(ur.Put(GinkgoT().Context(), &u)).To(Succeed()) + } + qc = quickconnect.New() + api = &Router{ds: ds, quickConnect: qc} + }) + + as := func(r *http.Request, u model.User) *http.Request { + return r.WithContext(request.WithUser(r.Context(), u)) + } + initiate := func() quickconnect.Request { + req, err := qc.Initiate(quickconnect.Device{ID: "dev-1", Name: "Pixel 7", App: "Finamp", AppVersion: "1.0.0"}) + Expect(err).ToNot(HaveOccurred()) + return req + } + + Describe("GET /QuickConnect/Enabled", func() { + DescribeTable("reports the config value", + func(enabled bool, expected string) { + conf.Server.Jellyfin.QuickConnect = enabled + w := httptest.NewRecorder() + api.quickConnectEnabled(w, httptest.NewRequest("GET", "/QuickConnect/Enabled", nil)) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Body.String()).To(MatchJSON(expected)) + }, + Entry("enabled", true, "true"), + Entry("disabled", false, "false"), + ) + }) + + Describe("requireQuickConnect", func() { + next := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusTeapot) }) + + It("rejects requests with 401 when Quick Connect is disabled", func() { + conf.Server.Jellyfin.QuickConnect = false + w := httptest.NewRecorder() + requireQuickConnect(next).ServeHTTP(w, httptest.NewRequest("POST", "/QuickConnect/Initiate", nil)) + Expect(w.Code).To(Equal(http.StatusUnauthorized)) + }) + + It("passes requests through when Quick Connect is enabled", func() { + w := httptest.NewRecorder() + requireQuickConnect(next).ServeHTTP(w, httptest.NewRequest("POST", "/QuickConnect/Initiate", nil)) + Expect(w.Code).To(Equal(http.StatusTeapot)) + }) + }) + + Describe("POST /QuickConnect/Initiate", func() { + initiateWith := func(authHeader string) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + r := httptest.NewRequest("POST", "/QuickConnect/Initiate", nil) + r.Header.Set("Authorization", authHeader) + api.quickConnectInitiate(w, r) + return w + } + + It("returns all QuickConnectResult fields, with the device from the auth header", func() { + w := initiateWith(finamp) + + Expect(w.Code).To(Equal(http.StatusOK)) + var raw map[string]any + Expect(json.Unmarshal(w.Body.Bytes(), &raw)).To(Succeed()) + // sdk-kotlin declares every field non-null. + Expect(raw).To(HaveKeyWithValue("Authenticated", false)) + Expect(raw).To(HaveKeyWithValue("Secret", MatchRegexp(`^[0-9a-f]{64}$`))) + Expect(raw).To(HaveKeyWithValue("Code", MatchRegexp(`^\d{6}$`))) + Expect(raw).To(HaveKeyWithValue("DeviceId", "dev-1")) + Expect(raw).To(HaveKeyWithValue("DeviceName", "Pixel 7")) + Expect(raw).To(HaveKeyWithValue("AppName", "Finamp")) + Expect(raw).To(HaveKeyWithValue("AppVersion", "1.0.0")) + Expect(raw).To(HaveKeyWithValue("DateAdded", Not(BeEmpty()))) + + _, err := qc.Status(raw["Secret"].(string)) + Expect(err).ToNot(HaveOccurred()) + }) + + It("returns 400 when the auth header does not identify the client", func() { + w := initiateWith(`MediaBrowser Client="Finamp", Device="Pixel 7", Version="1.0.0"`) + Expect(w.Code).To(Equal(http.StatusBadRequest)) + }) + + It("returns 400 when a client field is oversized", func() { + long := strings.Repeat("x", maxQuickConnectField+1) + api.quickConnect = fullQuickConnect{qc} + w := initiateWith(`MediaBrowser Client="Finamp", Device="` + long + `", DeviceId="dev-1", Version="1.0.0"`) + Expect(w.Code).To(Equal(http.StatusBadRequest)) + }) + + It("returns 429 when too many requests are pending", func() { + api.quickConnect = fullQuickConnect{qc} + Expect(initiateWith(finamp).Code).To(Equal(http.StatusTooManyRequests)) + }) + }) + + Describe("GET /QuickConnect/Connect", func() { + connect := func(secret string) (int, dto.QuickConnectResult) { + w := httptest.NewRecorder() + invoke(api.quickConnectConnect, w, httptest.NewRequest("GET", "/QuickConnect/Connect?Secret="+secret, nil)) + var res dto.QuickConnectResult + if w.Code == http.StatusOK { + Expect(json.Unmarshal(w.Body.Bytes(), &res)).To(Succeed()) + } + return w.Code, res + } + + It("reports whether the request has been authorized", func() { + req := initiate() + code, res := connect(req.Secret) + Expect(code).To(Equal(http.StatusOK)) + Expect(res.Authenticated).To(BeFalse()) + Expect(res.Code).To(Equal(req.Code)) + + _, err := qc.Authorize(req.Code, alice.ID) + Expect(err).ToNot(HaveOccurred()) + code, res = connect(req.Secret) + Expect(code).To(Equal(http.StatusOK)) + Expect(res.Authenticated).To(BeTrue()) + }) + + It("returns 404 for an unknown secret", func() { + code, _ := connect("unknown") + Expect(code).To(Equal(http.StatusNotFound)) + }) + }) + + Describe("POST /QuickConnect/Authorize", func() { + authorize := func(query string, u model.User) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + invoke(api.quickConnectAuthorize, w, as(httptest.NewRequest("POST", "/QuickConnect/Authorize?"+query, nil), u)) + return w + } + approver := func(req quickconnect.Request) string { + got, err := qc.Status(req.Secret) + Expect(err).ToNot(HaveOccurred()) + return got.UserID + } + + It("approves the code for the caller when no userId is given", func() { + req := initiate() + w := authorize("code="+req.Code, alice) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Body.String()).To(MatchJSON("true")) + Expect(approver(req)).To(Equal(alice.ID)) + }) + + It("accepts the caller's own userId", func() { + req := initiate() + w := authorize("Code="+req.Code+"&UserId="+dto.EncodeID(alice.ID), alice) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(approver(req)).To(Equal(alice.ID)) + }) + + It("forbids a non-admin from approving for another user", func() { + req := initiate() + w := authorize("code="+req.Code+"&userId="+dto.EncodeID(bob.ID), alice) + Expect(w.Code).To(Equal(http.StatusForbidden)) + Expect(approver(req)).To(BeEmpty()) + }) + + It("lets an admin approve for another user", func() { + req := initiate() + w := authorize("code="+req.Code+"&userId="+dto.EncodeID(bob.ID), admin) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(approver(req)).To(Equal(bob.ID)) + }) + + It("returns 404 when an admin approves for an unknown user", func() { + req := initiate() + w := authorize("code="+req.Code+"&userId="+dto.EncodeID(testID("ghost")), admin) + Expect(w.Code).To(Equal(http.StatusNotFound)) + Expect(approver(req)).To(BeEmpty()) + }) + + It("returns 400 for a malformed userId", func() { + req := initiate() + w := authorize("code="+req.Code+"&userId=not-a-guid", admin) + Expect(w.Code).To(Equal(http.StatusBadRequest)) + }) + + It("returns 404 for an unknown code", func() { + w := authorize("code=000000", alice) + Expect(w.Code).To(Equal(http.StatusNotFound)) + }) + + It("returns 409 for a code that is already approved", func() { + req := initiate() + Expect(authorize("code="+req.Code, alice).Code).To(Equal(http.StatusOK)) + Expect(authorize("code="+req.Code, bob).Code).To(Equal(http.StatusConflict)) + Expect(approver(req)).To(Equal(alice.ID)) + }) + }) + + Describe("POST /Users/AuthenticateWithQuickConnect", func() { + redeem := func(body string) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + r := httptest.NewRequest("POST", "/Users/AuthenticateWithQuickConnect", strings.NewReader(body)) + r.Header.Set("Authorization", finamp) + api.authenticateWithQuickConnect(w, r) + return w + } + redeemSecret := func(secret string) *httptest.ResponseRecorder { + return redeem(`{"Secret":"` + secret + `"}`) + } + + It("signs in the approving user once", func() { + req := initiate() + _, _ = qc.Authorize(req.Code, alice.ID) + + w := redeemSecret(req.Secret) + Expect(w.Code).To(Equal(http.StatusOK)) + var res dto.AuthenticationResult + Expect(json.Unmarshal(w.Body.Bytes(), &res)).To(Succeed()) + Expect(res.User.Name).To(Equal("alice")) + Expect(res.ServerId).ToNot(BeEmpty()) + Expect(res.SessionInfo).ToNot(BeNil()) + Expect(res.SessionInfo.DeviceId).To(Equal("dev-1")) + claims, err := auth.Validate(res.AccessToken) + Expect(err).ToNot(HaveOccurred()) + Expect(claims.Subject).To(Equal("alice")) + + Expect(redeemSecret(req.Secret).Code).To(Equal(http.StatusNotFound)) + }) + + It("accepts a camelCase body", func() { + req := initiate() + _, _ = qc.Authorize(req.Code, alice.ID) + Expect(redeem(`{"secret":"` + req.Secret + `"}`).Code).To(Equal(http.StatusOK)) + }) + + It("returns 404 while the request is not approved", func() { + req := initiate() + Expect(redeemSecret(req.Secret).Code).To(Equal(http.StatusNotFound)) + }) + + DescribeTable("returns 400 for a bad body", + func(body string) { + Expect(redeem(body).Code).To(Equal(http.StatusBadRequest)) + }, + Entry("not JSON", `nope`), + Entry("no secret", `{}`), + ) + + It("returns 401 when the approving user no longer exists", func() { + req := initiate() + _, _ = qc.Authorize(req.Code, testID("ghost")) + Expect(redeemSecret(req.Secret).Code).To(Equal(http.StatusUnauthorized)) + }) + + It("returns 500 when the user lookup fails", func() { + req := initiate() + _, _ = qc.Authorize(req.Code, alice.ID) + ds.User().(*tests.MockedUserRepo).Error = errors.New("db down") + Expect(redeemSecret(req.Secret).Code).To(Equal(http.StatusInternalServerError)) + }) + }) +}) diff --git a/server/jellyfin/routing_test.go b/server/jellyfin/routing_test.go index 62b52cfff..70da5b1eb 100644 --- a/server/jellyfin/routing_test.go +++ b/server/jellyfin/routing_test.go @@ -19,7 +19,7 @@ var _ = Describe("Case-insensitive routing", func() { var api *Router BeforeEach(func() { - api = New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + api = New(&tests.MockDataStore{}, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) }) It("serves a fully lowercase path directly", func() { diff --git a/server/jellyfin/sessions.go b/server/jellyfin/sessions.go index a145e7359..e60b1bf5b 100644 --- a/server/jellyfin/sessions.go +++ b/server/jellyfin/sessions.go @@ -113,8 +113,8 @@ func (api *Router) reportPlaybackStopped(w http.ResponseWriter, r *http.Request) w.WriteHeader(http.StatusNoContent) } -// postCapabilities acknowledges Jellyfin session-capability negotiation. -// Navidrome doesn't track per-session client capabilities, so this is a no-op. -func (api *Router) postCapabilities(w http.ResponseWriter, _ *http.Request) { +// acknowledge answers requests Navidrome keeps no state for: session capabilities (not tracked per +// session) and playback pings (transcodes live only as long as their stream request). +func (api *Router) acknowledge(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) } diff --git a/server/jellyfin/sessions_test.go b/server/jellyfin/sessions_test.go index 05151c39a..f2f04497c 100644 --- a/server/jellyfin/sessions_test.go +++ b/server/jellyfin/sessions_test.go @@ -152,12 +152,12 @@ var _ = Describe("Sessions", func() { }) }) - Describe("postCapabilities", func() { + Describe("acknowledge", func() { It("returns 204 No Content and does not touch the scrobbler", func() { w := httptest.NewRecorder() r := authed(httptest.NewRequest("POST", "/Sessions/Capabilities", strings.NewReader(`{"SupportsMediaControl":true}`))) - api.postCapabilities(w, r) + api.acknowledge(w, r) Expect(w.Code).To(Equal(http.StatusNoContent)) Expect(pt.reported).To(BeEmpty()) diff --git a/server/jellyfin/similar.go b/server/jellyfin/similar.go index 82d12f52b..50503160c 100644 --- a/server/jellyfin/similar.go +++ b/server/jellyfin/similar.go @@ -111,11 +111,26 @@ func (api *Router) getSimilarAlbums(w http.ResponseWriter, r *http.Request) { // a track seed leads its own mix; provider errors and unknown seeds degrade to seed-only/empty // results, never a 404 the client would surface as an error. func (api *Router) getInstantMix(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() id, ok := itemIDParam(w, r, "itemId") if !ok { return } + api.instantMix(w, r, id) +} + +// getInstantMixByQuery serves the legacy Artists/InstantMix and MusicGenres/InstantMix forms, +// which take the seed as ?id= instead of a path segment. +func (api *Router) getInstantMixByQuery(w http.ResponseWriter, r *http.Request) { + id, ok := dto.DecodeID(req.Params(r).StringOr("id", "")) + if !ok { + http.Error(w, "Not Found", http.StatusNotFound) + return + } + api.instantMix(w, r, id) +} + +func (api *Router) instantMix(w http.ResponseWriter, r *http.Request, id string) { + ctx := r.Context() limit := clampLimit(req.Params(r).IntOr("limit", 0), defaultSimilarLimit, maxInstantMixLimit) // Genre ids don't resolve via GetEntityByID, so a not-found entity is fine: it is just "not a @@ -201,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 401c791fd..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,14 +86,14 @@ 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()) token = t - api = New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + api = New(ds, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) }) It("upgrades when authenticated via the api_key query parameter", func() { diff --git a/server/jellyfin/stream.go b/server/jellyfin/stream.go index 87660056c..515a6fcb8 100644 --- a/server/jellyfin/stream.go +++ b/server/jellyfin/stream.go @@ -1,8 +1,10 @@ package jellyfin import ( + "cmp" "fmt" "math" + "mime" "net/http" "net/url" "slices" @@ -10,6 +12,7 @@ import ( "strings" "github.com/go-chi/chi/v5" + "github.com/navidrome/navidrome/core/stream" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" @@ -26,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 @@ -65,16 +68,14 @@ func (api *Router) getPlaybackInfo(w http.ResponseWriter, r *http.Request) { api.ok(w, r, dto.PlaybackInfoResponse{MediaSources: []dto.MediaSourceInfo{src}, PlaySessionId: dto.EncodeID(mf.ID)}) } -// streamAudio serves /Audio/{itemId}/stream[.container] and /Audio/{itemId}/universal, -// reusing the same transcode-decision + streaming pipeline as the Subsonic /stream endpoint. +// streamAudio serves /Audio/{itemId}/stream[.container], reusing the same transcode-decision + +// streaming pipeline as the Subsonic /stream endpoint. func (api *Router) streamAudio(w http.ResponseWriter, r *http.Request) { mf, ok := api.mediaFileForRequest(w, r) if !ok { return } - ctx := r.Context() p := req.Params(r) - format := p.StringOr("container", "") if format == "" { // The /stream.{container} route form carries the format as a path segment, not a query param. @@ -84,17 +85,69 @@ func (api *Router) streamAudio(w http.ResponseWriter, r *http.Request) { // Jellyfin's audioCodec param names the target codec when no container is given. format = p.StringOr("audiocodec", "") } + api.serveAudio(w, r, mf, format) +} + +// streamUniversal serves /Audio/{itemId}/universal, where Container lists the "container|codec" +// entries the client direct-plays, and TranscodingContainer/AudioCodec name the fallback target. +func (api *Router) streamUniversal(w http.ResponseWriter, r *http.Request) { + mf, ok := api.mediaFileForRequest(w, r) + if !ok { + return + } + p := req.Params(r) + var streamReq stream.Request if p.BoolOr("static", false) { + streamReq = api.transcodeDecider.ResolveRequest(r.Context(), mf, "raw", 0, 0) + } else { + streamReq = api.transcodeDecider.ResolveClientRequest(r.Context(), mf, universalClientInfo(p), 0) + } + api.serveStream(w, r, mf, streamReq) +} + +func universalClientInfo(p *req.Values) *stream.ClientInfo { + ci := &stream.ClientInfo{Name: "jellyfin-universal", MaxAudioBitrate: bitRateParam(p)} + for entry := range strings.SplitSeq(p.StringOr("container", ""), ",") { + container, codec, _ := strings.Cut(strings.TrimSpace(entry), "|") + if container == "" { + continue + } + profile := stream.DirectPlayProfile{Containers: []string{container}, Protocols: []string{stream.ProtocolHTTP}} + if codec != "" { + profile.AudioCodecs = []string{codec} + } + ci.DirectPlayProfiles = append(ci.DirectPlayProfiles, profile) + } + codec := p.StringOr("audiocodec", "") + if container := cmp.Or(p.StringOr("transcodingcontainer", ""), codec); container != "" { + ci.TranscodingProfiles = []stream.Profile{{Container: container, AudioCodec: cmp.Or(codec, container), Protocol: stream.ProtocolHTTP}} + ci.MaxTranscodingAudioBitrate = ci.MaxAudioBitrate + } + return ci +} + +// bitRateParam reads Jellyfin's bits/sec bitrate params as the kbps the stream package expects. +func bitRateParam(p *req.Values) int { + return cmp.Or(p.IntOr("audiobitrate", 0), p.IntOr("maxstreamingbitrate", 0)) / 1000 +} + +func (api *Router) serveAudio(w http.ResponseWriter, r *http.Request, mf *model.MediaFile, format string) { + if req.Params(r).BoolOr("static", false) { format = "raw" } + api.serveStream(w, r, mf, api.transcodeDecider.ResolveRequest(r.Context(), mf, format, bitRateParam(req.Params(r)), 0)) +} - // Bitrate params are bits/sec by Jellyfin convention; ResolveRequest expects kbps. - bitRate := p.IntOr("audiobitrate", 0) / 1000 - if bitRate == 0 { - bitRate = p.IntOr("maxstreamingbitrate", 0) / 1000 +func (api *Router) serveStream(w http.ResponseWriter, r *http.Request, mf *model.MediaFile, streamReq stream.Request) { + ctx := r.Context() + // A probe must not start a transcode. Omitting Content-Length is what tells a client (Fintunes) + // this is not direct play; direct play falls through to Serve, which answers HEAD itself. + if r.Method == http.MethodHead && streamReq.Format != "" && streamReq.Format != "raw" { + w.Header().Set("Content-Type", mime.TypeByExtension("."+streamReq.Format)) + w.Header().Set("Accept-Ranges", "none") + w.WriteHeader(http.StatusOK) + return } - - streamReq := api.transcodeDecider.ResolveRequest(ctx, mf, format, bitRate, 0) s, err := api.streamer.NewStream(ctx, mf, streamReq) if err != nil { api.internalError(w, r, err) @@ -160,15 +213,5 @@ func (api *Router) streamFile(w http.ResponseWriter, r *http.Request) { if !ok { return } - ctx := r.Context() - streamReq := api.transcodeDecider.ResolveRequest(ctx, mf, "raw", 0, 0) - s, err := api.streamer.NewStream(ctx, mf, streamReq) - if err != nil { - api.internalError(w, r, err) - return - } - defer s.Close() - if _, err := s.Serve(ctx, w, r); err != nil { - log.Error(ctx, "Jellyfin API: error streaming", "id", mf.ID, err) - } + api.serveStream(w, r, mf, api.transcodeDecider.ResolveRequest(r.Context(), mf, "raw", 0, 0)) } diff --git a/server/jellyfin/stream_test.go b/server/jellyfin/stream_test.go index d6220d33f..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") @@ -244,9 +244,78 @@ var _ = Describe("Stream", func() { }) }) + Describe("HEAD requests", func() { + head := func(query string) *httptest.ResponseRecorder { + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + {ID: testID("s1"), Title: "Song", Suffix: "flac", LibraryID: 1}, + }) + streamer.content = "audio-bytes" + w := httptest.NewRecorder() + r := httptest.NewRequest("HEAD", "/Audio/"+dto.EncodeID(testID("s1"))+"/stream?"+query, nil).WithContext(ctxUser()) + r = withChiURLParam(r, "itemId", dto.EncodeID(testID("s1"))) + invoke(api.streamAudio, w, r) + return w + } + + It("answers a transcode with the target type and no length, without starting it", func() { + w := head("audioCodec=mp3") + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(w.Header().Get("Content-Type")).To(Equal("audio/mpeg")) + Expect(w.Header().Get("Content-Length")).To(BeEmpty()) + Expect(w.Body.String()).To(BeEmpty()) + Expect(streamer.invoked).To(BeFalse()) + }) + + It("answers direct play through the streamer, without a body", func() { + w := head("static=true") + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(streamer.invoked).To(BeTrue()) + Expect(w.Body.String()).To(BeEmpty()) + }) + }) + + Describe("streamUniversal", func() { + universal := func(query string) { + ds.MediaFile().(*tests.MockMediaFileRepo).SetData(model.MediaFiles{ + {ID: testID("s1"), Suffix: "mp3", LibraryID: 1}, + }) + w := httptest.NewRecorder() + r := httptest.NewRequest("GET", "/Audio/"+dto.EncodeID(testID("s1"))+"/universal?"+query, nil).WithContext(ctxUser()) + r = withChiURLParam(r, "itemId", dto.EncodeID(testID("s1"))) + invoke(api.streamUniversal, w, r) + Expect(w.Code).To(Equal(http.StatusOK)) + } + + It("turns Container into direct play profiles and TranscodingContainer into the target", func() { + universal("Container=mp3,m4a|aac&TranscodingContainer=m4a&AudioCodec=aac&MaxStreamingBitrate=128000") + Expect(decider.client.DirectPlayProfiles).To(Equal([]stream.DirectPlayProfile{ + {Containers: []string{"mp3"}, Protocols: []string{stream.ProtocolHTTP}}, + {Containers: []string{"m4a"}, AudioCodecs: []string{"aac"}, Protocols: []string{stream.ProtocolHTTP}}, + })) + Expect(decider.client.TranscodingProfiles).To(Equal([]stream.Profile{ + {Container: "m4a", AudioCodec: "aac", Protocol: stream.ProtocolHTTP}, + })) + Expect(decider.client.MaxAudioBitrate).To(Equal(128)) + Expect(decider.client.MaxTranscodingAudioBitrate).To(Equal(128)) + }) + + It("uses AudioCodec as the target when no TranscodingContainer is given", func() { + universal("Container=mp3&AudioCodec=aac") + Expect(decider.client.TranscodingProfiles).To(Equal([]stream.Profile{ + {Container: "aac", AudioCodec: "aac", Protocol: stream.ProtocolHTTP}, + })) + }) + + It("serves the file as is for static=true", func() { + universal("static=true&Container=ogg") + Expect(decider.req.Format).To(Equal("raw")) + Expect(decider.client).To(BeNil()) + }) + }) + 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}, }) }) @@ -301,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" @@ -332,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() @@ -364,6 +433,7 @@ var _ = Describe("Stream", func() { type fakeTranscodeDecider struct { invoked bool req stream.Request + client *stream.ClientInfo } func (f *fakeTranscodeDecider) MakeDecision(context.Context, *model.MediaFile, *stream.ClientInfo, stream.TranscodeOptions) (*stream.TranscodeDecision, error) { @@ -384,6 +454,13 @@ func (f *fakeTranscodeDecider) ResolveRequest(_ context.Context, _ *model.MediaF return f.req } +func (f *fakeTranscodeDecider) ResolveClientRequest(_ context.Context, _ *model.MediaFile, ci *stream.ClientInfo, offset int) stream.Request { + f.invoked = true + f.client = ci + f.req = stream.Request{Offset: offset} + return f.req +} + // fakeMediaStreamer is a local test double for stream.MediaStreamer: it records whether // NewStream was invoked and, on success, returns a real (non-seekable) *stream.Stream backed // by an in-memory reader, so streamAudio's call to Stream.Serve exercises real code. diff --git a/server/jellyfin/system.go b/server/jellyfin/system.go index 8b847987f..2aae4c9c0 100644 --- a/server/jellyfin/system.go +++ b/server/jellyfin/system.go @@ -10,6 +10,7 @@ import ( "net/netip" "path" "strings" + "sync" "github.com/google/uuid" "github.com/navidrome/navidrome/conf" @@ -20,40 +21,41 @@ import ( "github.com/navidrome/navidrome/server/jellyfin/dto" ) -// jellyfinVersion is the Jellyfin API version advertised in the handshake. Clients feature-gate -// on it, so it must stay a real Jellyfin release, not Navidrome's own version. 10.9+ is required -// for Feishin to use the server lyrics endpoint. -const jellyfinVersion = "10.9.11" +// jellyfinVersion is the Jellyfin API version advertised in the handshake. Clients and SDKs gate on +// it (Streamyfin and the Android apps refuse < 10.10), and some parsers need exactly three parts. +const jellyfinVersion = "12.1.0" -func (api *Router) serverName() string { +func serverName() string { if conf.Server.Jellyfin.ServerName != "" { return conf.Server.Jellyfin.ServerName } return fmt.Sprintf("Navidrome %s", consts.Version) } -// serverID returns a stable Id that survives restarts, get-or-created in the Property table. -// Jellyfin clients cache ServerId across sessions, so a per-process value would break -// re-authentication. api.ds is nil only in unit tests; New() always sets it. -// -// The mutex serializes first-boot resolution so concurrent requests can't persist different -// UUIDs. Only a successful read or persisted id is cached; a transient failure yields a -// temporary id and retries on the next request rather than pinning a value. func (api *Router) serverID(ctx context.Context) string { - api.serverIDMu.Lock() - defer api.serverIDMu.Unlock() - if api.serverIDVal != "" { - return api.serverIDVal + return resolveServerID(ctx, api.ds, &api.serverIDVal) +} + +// Package-level: the Router and Discovery are separate objects and must not persist different ids. +var serverIDMu sync.Mutex + +// Clients cache ServerId across sessions, so it is persisted. Only a successful read or write is +// cached: a transient failure yields a temporary id and retries on the next call. +func resolveServerID(ctx context.Context, ds model.DataStore, cached *string) string { + serverIDMu.Lock() + defer serverIDMu.Unlock() + if *cached != "" { + return *cached } - if api.ds == nil { - api.serverIDVal = newServerID() - return api.serverIDVal + if ds == nil { + *cached = newServerID() + return *cached } - id, err := api.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 := api.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 } @@ -62,8 +64,8 @@ func (api *Router) serverID(ctx context.Context) string { return newServerID() } // Ids persisted before this change are dashed; normalize on read rather than rewriting the DB. - api.serverIDVal = strings.ReplaceAll(id, "-", "") - return api.serverIDVal + *cached = strings.ReplaceAll(id, "-", "") + return *cached } // newServerID returns a UUID in Jellyfin's no-dash GUID form (Guid.ToString("N")). @@ -75,7 +77,7 @@ func newServerID() string { func (api *Router) publicInfo(r *http.Request) dto.PublicSystemInfo { return dto.PublicSystemInfo{ LocalAddress: localAddress(r), - ServerName: api.serverName(), + ServerName: serverName(), Version: jellyfinVersion, ProductName: "Jellyfin Server", Id: api.serverID(r.Context()), @@ -107,7 +109,7 @@ func (api *Router) getSystemInfo(w http.ResponseWriter, r *http.Request) { func (api *Router) ping(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/plain; charset=utf-8") w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(api.serverName())) + _, _ = w.Write([]byte(serverName())) } // getEndpointInfo answers /System/Endpoint, which Finamp's connection test uses to pick between a @@ -136,7 +138,7 @@ func isSameMachine(r *http.Request, remote netip.Addr) bool { return parseIP(local.String()) == remote } -// remoteIP parses RemoteAddr, which the RealIP middleware may have rewritten to a bare IP. +// remoteIP parses RemoteAddr, which realIPMiddleware may have rewritten to a bare client IP. func remoteIP(r *http.Request) netip.Addr { return parseIP(r.RemoteAddr) } @@ -151,7 +153,3 @@ func parseIP(addr string) netip.Addr { } return ip.Unmap() } - -func (api *Router) quickConnectEnabled(w http.ResponseWriter, r *http.Request) { - api.ok(w, r, false) -} diff --git a/server/jellyfin/system_test.go b/server/jellyfin/system_test.go index 4b3ac240e..d339043a7 100644 --- a/server/jellyfin/system_test.go +++ b/server/jellyfin/system_test.go @@ -7,6 +7,7 @@ import ( "net" "net/http" "net/http/httptest" + "sync" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" @@ -144,17 +145,6 @@ var _ = Describe("System", func() { Expect(info.IsInNetwork).To(BeTrue()) }) - It("reports quick connect as disabled", func() { - w := httptest.NewRecorder() - r := httptest.NewRequest("GET", "/QuickConnect/Enabled", nil) - api.quickConnectEnabled(w, r) - - Expect(w.Code).To(Equal(http.StatusOK)) - var enabled bool - Expect(json.Unmarshal(w.Body.Bytes(), &enabled)).To(Succeed()) - Expect(enabled).To(BeFalse()) - }) - Context("serverID with a real DataStore", func() { var ctx context.Context var ds *tests.MockDataStore @@ -173,6 +163,17 @@ var _ = Describe("System", func() { Expect(second.serverID(ctx)).To(Equal(id)) }) + It("resolves one id when a Router and a Discovery race on first boot", func() { + r, d := &Router{ds: ds}, NewDiscovery(ds) + ids := make([]string, 2) + var wg sync.WaitGroup + wg.Go(func() { ids[0] = r.serverID(ctx) }) + wg.Go(func() { ids[1] = d.serverID(ctx) }) + wg.Wait() + Expect(ids[0]).ToNot(BeEmpty()) + Expect(ids[1]).To(Equal(ids[0])) + }) + It("memoizes the id across repeated calls on the same Router", func() { r := &Router{ds: ds} id := r.serverID(ctx) @@ -181,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()) @@ -193,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")) }) @@ -204,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 bbc60c892..d0a54ee63 100644 --- a/server/jellyfin/users.go +++ b/server/jellyfin/users.go @@ -13,10 +13,20 @@ import ( // getUserViews returns one CollectionFolder view per accessible library, so clients browse each // library as its own top-level view rather than one aggregate. func (api *Router) getUserViews(w http.ResponseWriter, r *http.Request) { - u, _ := request.UserFrom(r.Context()) - views := make([]dto.BaseItemDto, 0, len(u.Libraries)) - for _, lib := range u.Libraries { - views = append(views, libraryView(lib)) + ctx := r.Context() + 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().GetAll(ctx) + if err != nil { + api.internalError(w, r, err) + return + } + views := make([]dto.BaseItemDto, 0, len(libs)) + for _, lib := range libs { + if u.HasLibraryAccess(lib.ID) { + views = append(views, dto.LibraryToBaseItem(lib)) + } } api.ok(w, r, dto.QueryResult{Items: views, TotalRecordCount: len(views), StartIndex: 0}) } @@ -24,7 +34,7 @@ func (api *Router) getUserViews(w http.ResponseWriter, r *http.Request) { func (api *Router) getCurrentUser(w http.ResponseWriter, r *http.Request) { ctx := r.Context() u, _ := request.UserFrom(ctx) - api.ok(w, r, userToDto(&u, api.serverName(), api.serverID(ctx))) + api.ok(w, r, userToDto(&u, serverName(), api.serverID(ctx))) } // getPublicUsers advertises the users named in Jellyfin.ExposedPublicUsers for a client login @@ -45,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 3e53d598f..793a64a29 100644 --- a/server/jellyfin/users_test.go +++ b/server/jellyfin/users_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "time" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" @@ -17,12 +18,22 @@ 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 { - ctx := request.WithUser(context.Background(), model.User{ID: testID("u1"), UserName: "alice", Libraries: 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} + } + 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() { @@ -45,6 +56,25 @@ var _ = Describe("Users", func() { Expect(res.Items[1].Name).To(Equal("Podcasts")) }) + // Manet keeps no library, and so syncs no artists/albums/tracks, when these are missing. + It("describes the library like Jellyfin's CollectionFolder", func() { + created := time.Date(2026, 7, 1, 12, 0, 0, 0, time.UTC) + libs := model.Libraries{{ID: 1, Name: "Music", Path: "/music", TotalAlbums: 42, CreatedAt: created}} + w := httptest.NewRecorder() + api.getUserViews(w, authedWithLibraries(httptest.NewRequest("GET", "/UserViews", nil), libs)) + var res dto.QueryResult + Expect(json.Unmarshal(w.Body.Bytes(), &res)).To(Succeed()) + + item := res.Items[0] + Expect(*item.ChildCount).To(Equal(42)) + Expect(item.DateCreated).To(Equal("2026-07-01T12:00:00.0000000Z")) + Expect(item.SortName).To(Equal("Music")) + Expect(item.Path).To(Equal("/music")) + Expect(item.LocationType).To(Equal("FileSystem")) + Expect(item.UserData).ToNot(BeNil()) + Expect(item.UserData.ItemId).To(Equal(dto.EncodeLibraryID(1))) + }) + It("returns a single view for a user with one library", func() { libs := model.Libraries{{ID: 1, Name: "Music"}} w := httptest.NewRecorder() @@ -91,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 b710b4068..b65a2d6e1 100644 --- a/server/middlewares.go +++ b/server/middlewares.go @@ -7,7 +7,9 @@ import ( "errors" "fmt" "io/fs" + "net" "net/http" + "net/netip" "net/url" "strings" "time" @@ -15,6 +17,7 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5/middleware" "github.com/go-chi/cors" + "github.com/go-chi/httprate" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/log" @@ -165,20 +168,82 @@ func clientUniqueIDMiddleware(next http.Handler) http.Handler { }) } -// realIPMiddleware applies middleware.RealIP, and additionally saves the request's original RemoteAddr to the request's -// context if navidrome is behind a trusted reverse proxy. +// realIPMiddleware resolves the request's client IP into the context, where it can be read with +// middleware.GetClientIP, and mirrors it into RemoteAddr for logging and player registration. +// Forwarding headers are only honoured when the peer is listed in ExtAuth.TrustedSources, so that +// a client cannot pick its own identity and evade controls keyed on it. The peer address is kept +// in the context as request.ReverseProxyIp. func realIPMiddleware(next http.Handler) http.Handler { - if conf.Server.ExtAuth.TrustedSources != "" { - return chi.Chain( - reqToCtx(request.ReverseProxyIp, func(r *http.Request) any { return r.RemoteAddr }), - middleware.RealIP, - ).Handler(next) + trusted := conf.Server.ExtAuth.TrustedSources + fromPeer := middleware.ClientIPFromRemoteAddr(next) + if trusted == "" { + return fromPeer } - // The middleware is applied without a trusted reverse proxy to support other use-cases such as multiple clients - // behind a caching proxy. In this case, navidrome only uses the request's RemoteAddr for logging, so the security - // impact of reading the headers from untrusted sources is limited. - return middleware.RealIP(next) + // Last match wins, so this order reproduces RealIP's precedence: True-Client-IP, X-Real-IP, + // X-Forwarded-For, peer. Only X-Forwarded-For is checked against the trusted list. + fromProxy := chi.Chain( + middleware.ClientIPFromRemoteAddr, + middleware.ClientIPFromXFF(trustedProxyPrefixes(trusted)...), + middleware.ClientIPFromHeader("X-Real-IP"), + middleware.ClientIPFromHeader("True-Client-IP"), + ).Handler(mirrorClientIP(next)) + + dispatch := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if validateIPAgainstList(r.RemoteAddr, trusted) { + fromProxy.ServeHTTP(w, r) + return + } + log.Trace(r.Context(), "Ignoring forwarding headers from untrusted peer", "peer", r.RemoteAddr) + fromPeer.ServeHTTP(w, r) + }) + return reqToCtx(request.ReverseProxyIp, func(r *http.Request) any { return r.RemoteAddr })(dispatch) +} + +// mirrorClientIP copies the resolved client IP into RemoteAddr when it differs from the peer. +func mirrorClientIP(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if ip := middleware.GetClientIP(r.Context()); ip != "" && ip != peerHost(r) { + r.RemoteAddr = ip + } + next.ServeHTTP(w, r) + }) +} + +// peerHost returns the host part of RemoteAddr, which may already be a bare IP. +func peerHost(r *http.Request) string { + if host, _, err := net.SplitHostPort(r.RemoteAddr); err == nil { + return host + } + return r.RemoteAddr +} + +// trustedProxyPrefixes returns the CIDR entries of a trusted sources list, skipping non-CIDR +// entries such as the "@" unix socket marker. An empty result makes ClientIPFromXFF trust +// exactly one hop. +func trustedProxyPrefixes(list string) []string { + var prefixes []string + for _, entry := range strings.Split(list, ",") { + entry = strings.TrimSpace(entry) + if _, err := netip.ParsePrefix(entry); err == nil { + prefixes = append(prefixes, entry) + } + } + return prefixes +} + +// ClientIPRateLimiter returns a rate limiter keyed by ClientIP, so spoofed forwarding headers +// cannot be rotated for a fresh bucket. +func ClientIPRateLimiter(requestLimit int, windowLength time.Duration) func(http.Handler) http.Handler { + return httprate.LimitBy(requestLimit, windowLength, func(r *http.Request) (string, error) { + return ClientIP(r), nil + }) +} + +// ClientIP returns the canonical client IP resolved by realIPMiddleware, for keying rate limits. The +// peer address fallback degrades a missing middleware to per-peer limiting, not one shared bucket. +func ClientIP(r *http.Request) string { + return httprate.CanonicalizeIP(cmp.Or(middleware.GetClientIP(r.Context()), peerHost(r))) } // reqToCtx creates a middleware that updates the request's context with a value computed from the request. A given key @@ -315,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 4ba9a853b..15cf70341 100644 --- a/server/middlewares_test.go +++ b/server/middlewares_test.go @@ -9,6 +9,7 @@ import ( "time" "github.com/go-chi/chi/v5" + "github.com/go-chi/chi/v5/middleware" "github.com/google/uuid" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/configtest" @@ -380,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) @@ -406,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 @@ -421,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)) }) }) @@ -430,9 +431,105 @@ 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)) }) }) }) + Describe("realIPMiddleware", func() { + var resolved, remoteAddr string + var proxyIP any + next := func(w http.ResponseWriter, r *http.Request) { + resolved = middleware.GetClientIP(r.Context()) + remoteAddr = r.RemoteAddr + proxyIP = r.Context().Value(request.ReverseProxyIp) + } + call := func(peer string, headers map[string]string) { + resolved, remoteAddr, proxyIP = "", "", nil + r := httptest.NewRequest("POST", "/auth/login", nil) + r.RemoteAddr = peer + for k, v := range headers { + r.Header.Set(k, v) + } + realIPMiddleware(http.HandlerFunc(next)).ServeHTTP(httptest.NewRecorder(), r) + } + + Context("without a trusted proxy", func() { + It("ignores client-supplied forwarding headers", func() { + call("10.0.0.1:1234", map[string]string{ + "X-Forwarded-For": "203.0.113.5", + "X-Real-IP": "203.0.113.6", + "True-Client-IP": "203.0.113.7", + }) + Expect(resolved).To(Equal("10.0.0.1")) + }) + It("leaves RemoteAddr untouched", func() { + call("10.0.0.1:1234", map[string]string{"X-Forwarded-For": "203.0.113.5"}) + Expect(remoteAddr).To(Equal("10.0.0.1:1234")) + }) + }) + + Context("with a trusted proxy", func() { + BeforeEach(func() { + conf.Server.ExtAuth.TrustedSources = "10.0.0.0/8" + }) + It("uses the forwarded client IP when the peer is a trusted proxy", func() { + call("10.0.0.1:1234", map[string]string{"X-Forwarded-For": "203.0.113.5, 10.0.0.1"}) + Expect(resolved).To(Equal("203.0.113.5")) + Expect(remoteAddr).To(Equal("203.0.113.5")) + }) + It("honours X-Real-IP from a trusted proxy", func() { + call("10.0.0.1:1234", map[string]string{"X-Real-IP": "203.0.113.6"}) + Expect(resolved).To(Equal("203.0.113.6")) + }) + It("ignores forwarding headers when the peer is not a trusted proxy", func() { + call("198.51.100.9:1234", map[string]string{"X-Forwarded-For": "203.0.113.5"}) + Expect(resolved).To(Equal("198.51.100.9")) + Expect(remoteAddr).To(Equal("198.51.100.9:1234")) + }) + It("keeps the peer address in the context for external auth", func() { + call("10.0.0.1:1234", map[string]string{"X-Forwarded-For": "203.0.113.5"}) + Expect(proxyIP).To(Equal("10.0.0.1:1234")) + }) + }) + }) + + Describe("ClientIPRateLimiter", func() { + var handler http.Handler + JustBeforeEach(func() { + handler = realIPMiddleware(ClientIPRateLimiter(2, time.Minute)( + http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }))) + }) + attempt := func(peer string, header, value string) int { + r := httptest.NewRequest("POST", "/auth/login", nil) + r.RemoteAddr = peer + r.Header.Set(header, value) + w := httptest.NewRecorder() + handler.ServeHTTP(w, r) + return w.Code + } + + DescribeTable("keeps one bucket per peer when the forwarding header is rotated", + func(header string) { + Expect(attempt("198.51.100.9:1", header, "203.0.113.1")).To(Equal(http.StatusOK)) + Expect(attempt("198.51.100.9:2", header, "203.0.113.2")).To(Equal(http.StatusOK)) + Expect(attempt("198.51.100.9:3", header, "203.0.113.3")).To(Equal(http.StatusTooManyRequests)) + }, + Entry("X-Forwarded-For", "X-Forwarded-For"), + Entry("X-Real-IP", "X-Real-IP"), + Entry("True-Client-IP", "True-Client-IP"), + ) + + Context("behind a trusted proxy", func() { + BeforeEach(func() { + conf.Server.ExtAuth.TrustedSources = "10.0.0.0/8" + }) + It("gives each real client its own bucket", func() { + Expect(attempt("10.0.0.1:1", "X-Forwarded-For", "203.0.113.1")).To(Equal(http.StatusOK)) + Expect(attempt("10.0.0.1:2", "X-Forwarded-For", "203.0.113.1")).To(Equal(http.StatusOK)) + Expect(attempt("10.0.0.1:3", "X-Forwarded-For", "203.0.113.1")).To(Equal(http.StatusTooManyRequests)) + Expect(attempt("10.0.0.1:4", "X-Forwarded-For", "203.0.113.2")).To(Equal(http.StatusOK)) + }) + }) + }) }) 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 d1007f457..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) + 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/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/inspect.go b/server/nativeapi/inspect.go index 7c96312ed..61013dd5a 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 } @@ -22,7 +22,12 @@ func doInspect(ctx context.Context, ds model.DataStore, id string) (*core.Inspec return nil, model.ErrNotFound } - return core.Inspect(file.AbsolutePath(), file.LibraryID, file.FolderID) + lib, err := ds.Library().Get(ctx, file.LibraryID) + if err != nil { + return nil, err + } + + return core.Inspect(file.AbsolutePath(), *lib, file.FolderID) } func inspect(ds model.DataStore) http.HandlerFunc { diff --git a/server/nativeapi/library_test.go b/server/nativeapi/library_test.go index 13b33c238..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) + 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 ebc9aeb28..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) + 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 0ad9bb0cc..a906e36b9 100644 --- a/server/nativeapi/missing.go +++ b/server/nativeapi/missing.go @@ -6,7 +6,6 @@ import ( "maps" "net/http" - "github.com/Masterminds/squirrel" "github.com/deluan/rest" "github.com/navidrome/navidrome/core" "github.com/navidrome/navidrome/log" @@ -15,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 { @@ -45,22 +41,15 @@ func (r *missingRepository) parseOptions(options []rest.QueryOptions) rest.Query return opt } -func (r *missingRepository) Read(id string) (any, error) { - all, err := r.mfRepo.GetAll(model.QueryOptions{Filters: squirrel.And{ - squirrel.Eq{"id": id}, - squirrel.Eq{"missing": true}, - }}) +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 } - if len(all) == 0 { + if !mf.Missing { return nil, model.ErrNotFound } - return all[0], nil -} - -func (r *missingRepository) EntityName() string { - return "missing_files" + return mf, nil } func deleteMissingFiles(maintenance core.Maintenance) http.HandlerFunc { @@ -90,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 new file mode 100644 index 000000000..a53f13fc3 --- /dev/null +++ b/server/nativeapi/missing_test.go @@ -0,0 +1,57 @@ +package nativeapi + +import ( + "bytes" + "net/http" + "net/http/httptest" + "time" + + "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" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Missing Files Endpoint", func() { + var router http.Handler + var token string + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.SessionTimeout = time.Minute + conf.Server.EnableSharing = false + + mfRepo := tests.CreateMockMediaFileRepo() + mfRepo.SetData(model.MediaFiles{ + {ID: "missing-1", Title: "Gone", Missing: true}, + {ID: "present-1", Title: "Here"}, + }) + userRepo := tests.CreateMockUserRepo() + ds := &tests.MockDataStore{MockedMediaFile: mfRepo, MockedUser: userRepo, MockedProperty: &tests.MockedPropertyRepo{}} + auth.Init(ds) + + user := model.User{ID: "user-1", UserName: "user", NewPassword: "pass"} + 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, playlists.NewPlaylists(ds, nil), nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil)) + }) + + DescribeTable("GET /missing/{id}", + func(id string, status int) { + w := httptest.NewRecorder() + router.ServeHTTP(w, createAuthenticatedRequest(http.MethodGet, "/missing/"+id, &bytes.Buffer{}, token)) + Expect(w.Code).To(Equal(status), w.Body.String()) + }, + Entry("returns a missing file", "missing-1", http.StatusOK), + Entry("returns 404 for a file that is not missing", "present-1", http.StatusNotFound), + Entry("returns 404 for an unknown id", "unknown", http.StatusNotFound), + ) +}) diff --git a/server/nativeapi/native_api.go b/server/nativeapi/native_api.go index 57a712a20..97ad14be2 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" @@ -17,6 +15,7 @@ import ( "github.com/navidrome/navidrome/core/external" "github.com/navidrome/navidrome/core/metrics" playlistsvc "github.com/navidrome/navidrome/core/playlists" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/model/request" @@ -48,10 +47,11 @@ type Router struct { pluginManager PluginManager imgUpload artwork.Uploader provider external.Provider + quickConnect quickconnect.QuickConnect } -func New(ds model.DataStore, share core.Share, playlists playlistsvc.Playlists, insights metrics.Insights, libraryService core.Library, userService core.User, maintenance core.Maintenance, pluginManager PluginManager, imgUpload artwork.Uploader, provider external.Provider) *Router { - r := &Router{ds: ds, share: share, playlists: playlists, insights: insights, libs: libraryService, users: userService, maintenance: maintenance, pluginManager: pluginManager, imgUpload: imgUpload, provider: provider} +func New(ds model.DataStore, share core.Share, playlists playlistsvc.Playlists, insights metrics.Insights, libraryService core.Library, userService core.User, maintenance core.Maintenance, pluginManager PluginManager, imgUpload artwork.Uploader, provider external.Provider, quickConnect quickconnect.QuickConnect) *Router { + r := &Router{ds: ds, share: share, playlists: playlists, insights: insights, libs: libraryService, users: userService, maintenance: maintenance, pluginManager: pluginManager, imgUpload: imgUpload, provider: provider, quickConnect: quickConnect} r.Handler = r.routes() return r } @@ -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) @@ -88,6 +88,7 @@ func (api *Router) routes() http.Handler { api.addMissingFilesRoute(r) api.addKeepAliveRoute(r) api.addInsightsRoute(r) + api.addQuickConnectRoute(r) r.With(adminOnlyMiddleware).Group(func(r chi.Router) { api.addInspectRoute(r) @@ -95,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) @@ -143,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)) }) @@ -197,28 +189,24 @@ 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)) }) } 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/nativeapi/native_api_song_test.go b/server/nativeapi/native_api_song_test.go index 203fcd4cf..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) + 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 d215eb9dd..82f138492 100644 --- a/server/nativeapi/playlists.go +++ b/server/nativeapi/playlists.go @@ -16,10 +16,9 @@ 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 - // writePlaylistError maps a playlist service error to an HTTP status, or defaultStatus if unknown. func writePlaylistError(w http.ResponseWriter, err error, defaultStatus int) { switch { @@ -34,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)) @@ -42,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) { @@ -60,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 { @@ -101,8 +100,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 9abcc477f..82e3bc86a 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" @@ -96,10 +97,10 @@ 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) + nativeRouter := New(ds, nil, plsSvc, nil, tests.NewMockLibraryService(), tests.NewMockUserService(), nil, nil, nil, nil, nil) router = server.JWTVerifier(nativeRouter) w = httptest.NewRecorder() }) @@ -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) { @@ -192,6 +224,7 @@ var _ = Describe("writePlaylistError", func() { }, Entry("not found -> 404", model.ErrNotFound, http.StatusNotFound), Entry("not authorized -> 403", model.ErrNotAuthorized, http.StatusForbidden), + Entry("rest permission denied -> 403", rest.ErrPermissionDenied, http.StatusForbidden), Entry("not editable -> 409", model.ErrPlaylistNotEditable, http.StatusConflict), Entry("unrecognized -> default", model.ErrValidation, http.StatusBadRequest), ) @@ -202,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 @@ -229,7 +254,9 @@ 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 } @@ -248,6 +275,17 @@ func (m *mockPlaylistsService) SetImage(ctx context.Context, id string, reader i return model.ErrNotFound } -func (m *mockPlaylistsService) TracksRepository(_ context.Context, _ string, _ bool) rest.Repository { +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 + } + return m.playlist, nil +} + +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 1683885e7..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) + 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/quickconnect.go b/server/nativeapi/quickconnect.go new file mode 100644 index 000000000..61d402cc5 --- /dev/null +++ b/server/nativeapi/quickconnect.go @@ -0,0 +1,76 @@ +package nativeapi + +import ( + "encoding/json" + "errors" + "net/http" + + "github.com/go-chi/chi/v5" + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/core/quickconnect" + "github.com/navidrome/navidrome/log" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/server" +) + +type quickConnectDevice struct { + AppName string `json:"appName"` + AppVersion string `json:"appVersion"` + DeviceName string `json:"deviceName"` +} + +func (api *Router) addQuickConnectRoute(r chi.Router) { + if !quickconnect.Enabled() { + return + } + r.Route("/quickconnect", func(r chi.Router) { + // Throttled like login so a signed-in user cannot enumerate other people's pending codes. + if conf.Server.AuthRequestLimit > 0 { + r.Use(server.ClientIPRateLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength)) + } + r.Get("/", lookupQuickConnect(api.quickConnect)) + r.Post("/authorize", authorizeQuickConnect(api.quickConnect)) + }) +} + +func lookupQuickConnect(qc quickconnect.QuickConnect) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + req, err := qc.Lookup(r.URL.Query().Get("code")) + writeQuickConnectResult(w, r, req, err) + } +} + +func authorizeQuickConnect(qc quickconnect.QuickConnect) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var body struct { + Code string `json:"code"` + } + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + http.Error(w, "Bad Request", http.StatusBadRequest) + return + } + ctx := r.Context() + user, _ := request.UserFrom(ctx) + req, err := qc.Authorize(body.Code, user.ID) + if err == nil { + log.Info(ctx, "Quick Connect sign-in approved", "username", user.UserName, "client", req.Device.App, "device", req.Device.Name) + } + writeQuickConnectResult(w, r, req, err) + } +} + +func writeQuickConnectResult(w http.ResponseWriter, r *http.Request, req quickconnect.Request, err error) { + switch { + case errors.Is(err, model.ErrNotFound): + http.Error(w, "Unknown code", http.StatusNotFound) + case errors.Is(err, quickconnect.ErrAlreadyAuthorized): + http.Error(w, "Code already used", http.StatusConflict) + case err != nil: + log.Error(r.Context(), "Quick Connect failed", err) + http.Error(w, "Internal Server Error", http.StatusInternalServerError) + default: + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(quickConnectDevice{AppName: req.Device.App, AppVersion: req.Device.AppVersion, DeviceName: req.Device.Name}) + } +} diff --git a/server/nativeapi/quickconnect_test.go b/server/nativeapi/quickconnect_test.go new file mode 100644 index 000000000..f2f06f0f3 --- /dev/null +++ b/server/nativeapi/quickconnect_test.go @@ -0,0 +1,148 @@ +package nativeapi + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "time" + + "github.com/go-chi/chi/v5" + "github.com/navidrome/navidrome/conf" + "github.com/navidrome/navidrome/conf/configtest" + "github.com/navidrome/navidrome/core/quickconnect" + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Quick Connect endpoints", func() { + var ( + qc quickconnect.QuickConnect + pending quickconnect.Request + user = model.User{ID: "u1", UserName: "alice"} + ) + + BeforeEach(func() { + qc = quickconnect.New() + var err error + pending, err = qc.Initiate(quickconnect.Device{ID: "dev-1", Name: "Pixel 7", App: "Finamp", AppVersion: "1.0.0"}) + Expect(err).ToNot(HaveOccurred()) + }) + + asUser := func(r *http.Request) *http.Request { + return r.WithContext(request.WithUser(r.Context(), user)) + } + decode := func(w *httptest.ResponseRecorder) quickConnectDevice { + var res quickConnectDevice + Expect(json.Unmarshal(w.Body.Bytes(), &res)).To(Succeed()) + return res + } + finamp := quickConnectDevice{AppName: "Finamp", AppVersion: "1.0.0", DeviceName: "Pixel 7"} + + Describe("GET /quickconnect", func() { + lookup := func(code string) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + lookupQuickConnect(qc)(w, asUser(httptest.NewRequest("GET", "/quickconnect?code="+code, nil))) + return w + } + + It("describes the device waiting for the code", func() { + w := lookup(pending.Code) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(decode(w)).To(Equal(finamp)) + }) + + It("does not approve the code", func() { + lookup(pending.Code) + got, _ := qc.Status(pending.Secret) + Expect(got.Authorized()).To(BeFalse()) + }) + + It("returns 404 for an unknown code", func() { + Expect(lookup("000000").Code).To(Equal(http.StatusNotFound)) + }) + + It("returns 409 for a code that is already approved", func() { + _, _ = qc.Authorize(pending.Code, "someone") + Expect(lookup(pending.Code).Code).To(Equal(http.StatusConflict)) + }) + }) + + Describe("POST /quickconnect/authorize", func() { + authorize := func(body string) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + authorizeQuickConnect(qc)(w, asUser(httptest.NewRequest("POST", "/quickconnect/authorize", strings.NewReader(body)))) + return w + } + + It("approves the code for the signed-in user", func() { + w := authorize(`{"code":"` + pending.Code + `"}`) + Expect(w.Code).To(Equal(http.StatusOK)) + Expect(decode(w)).To(Equal(finamp)) + got, _ := qc.Status(pending.Secret) + Expect(got.UserID).To(Equal(user.ID)) + }) + + It("returns 404 for an unknown code", func() { + Expect(authorize(`{"code":"000000"}`).Code).To(Equal(http.StatusNotFound)) + }) + + It("returns 409 for a code that is already approved", func() { + _, _ = qc.Authorize(pending.Code, "someone") + Expect(authorize(`{"code":"` + pending.Code + `"}`).Code).To(Equal(http.StatusConflict)) + }) + + It("returns 400 for a bad body", func() { + Expect(authorize(`nope`).Code).To(Equal(http.StatusBadRequest)) + }) + }) + + Describe("route registration", func() { + serve := func() int { + r := chi.NewRouter() + (&Router{quickConnect: qc}).addQuickConnectRoute(r) + w := httptest.NewRecorder() + r.ServeHTTP(w, asUser(httptest.NewRequest("GET", "/quickconnect?code="+pending.Code, nil))) + return w.Code + } + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.Jellyfin.Enabled = true + conf.Server.Jellyfin.QuickConnect = true + }) + + It("mounts the routes when Jellyfin Quick Connect is on", func() { + Expect(serve()).To(Equal(http.StatusOK)) + }) + + It("rate-limits code attempts when a login limit is configured", func() { + conf.Server.AuthRequestLimit = 2 + conf.Server.AuthWindowLength = time.Minute + r := chi.NewRouter() + (&Router{quickConnect: qc}).addQuickConnectRoute(r) + attempt := func() int { + w := httptest.NewRecorder() + req := httptest.NewRequest("GET", "/quickconnect?code=000000", nil) + req.RemoteAddr = "10.0.0.1:1234" + r.ServeHTTP(w, asUser(req)) + return w.Code + } + Expect(attempt()).To(Equal(http.StatusNotFound)) + Expect(attempt()).To(Equal(http.StatusNotFound)) + Expect(attempt()).To(Equal(http.StatusTooManyRequests)) + }) + + It("skips the routes when the Jellyfin API is off", func() { + conf.Server.Jellyfin.Enabled = false + Expect(serve()).To(Equal(http.StatusNotFound)) + }) + + It("skips the routes when Quick Connect is off", func() { + conf.Server.Jellyfin.QuickConnect = false + Expect(serve()).To(Equal(http.StatusNotFound)) + }) + }) +}) 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 2a363980f..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) + 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_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/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..4de2a75c1 --- /dev/null +++ b/server/public/handle_shares_test.go @@ -0,0 +1,55 @@ +package public + +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" + . "github.com/onsi/gomega" +) + +var _ = Describe("handleM3U", func() { + var ds *tests.MockDataStore + var shareRepo *tests.MockShareRepo + var pub *Router + + BeforeEach(func() { + auth.PublicTokenAuth = jwtauth.New("HS256", []byte("test-secret"), nil) + 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)) + }) +}) diff --git a/server/public/handle_streams.go b/server/public/handle_streams.go index 3d624f661..46c7ca210 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) @@ -60,9 +60,8 @@ func (pub *Router) handleStream(w http.ResponseWriter, r *http.Request) { return } - stream, err := pub.streamer.NewStream(ctx, mf, streampkg.Request{ - Format: info.format, BitRate: info.bitrate, - }) + streamReq := pub.decider.ResolveRequest(ctx, mf, info.format, info.bitrate, 0) + stream, err := pub.streamer.NewStream(ctx, mf, streamReq) if err != nil { if errors.Is(err, streampkg.ErrTooManyTranscodes) { w.Header().Set("Retry-After", strconv.Itoa(streampkg.RetryAfterSeconds)) diff --git a/server/public/handle_streams_test.go b/server/public/handle_streams_test.go index 870dfa8ef..965bc7e05 100644 --- a/server/public/handle_streams_test.go +++ b/server/public/handle_streams_test.go @@ -107,18 +107,20 @@ 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{} + ds = &tests.MockDataStore{MockedTranscoding: &tests.MockTranscodingRepo{}} shareRepo = &tests.MockShareRepo{} ds.MockedShare = shareRepo streamer = &mockStreamer{} - pub = &Router{ds: ds, streamer: streamer} + pub = &Router{ds: ds, streamer: streamer, decider: stream.NewTranscodeDecider(ds, tests.NewMockFFmpeg(""))} }) makeRequest := func(token string) *httptest.ResponseRecorder { @@ -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}) @@ -150,8 +152,17 @@ var _ = Describe("handleStream", func() { makeRequest(token) Expect(streamer.called).To(BeTrue()) - Expect(streamer.req.Format).To(Equal("mp3")) - Expect(streamer.req.BitRate).To(Equal(192)) + }) + + It("resolves the full stream request like the Subsonic endpoint, so transcodes share the cache", func() { + mf := model.MediaFile{ID: "mf-123", Suffix: "flac", BitRate: 1500, SampleRate: 44100, BitDepth: new(24), Channels: 2} + shareOwnedBy(model.User{ID: "owner1", UserName: "owner1", IsAdmin: true}, mf) + + claims := auth.Claims{ID: "mf-123", Format: "opus", BitRate: 128, ShareID: "share123"} + token, _ := auth.CreateExpiringPublicToken(time.Now().Add(time.Hour), claims) + makeRequest(token) + + Expect(streamer.req).To(Equal(stream.Request{Format: "opus", BitRate: 128, SampleRate: 48000, Channels: 2})) }) It("returns 404 when the track is outside the share owner's libraries", func() { @@ -171,7 +182,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/public/public.go b/server/public/public.go index 142c474bd..8239ef927 100644 --- a/server/public/public.go +++ b/server/public/public.go @@ -21,14 +21,15 @@ type Router struct { http.Handler artwork artwork.Artwork streamer stream.MediaStreamer + decider stream.TranscodeDecider archiver core.Archiver share core.Share assetsHandler http.Handler ds model.DataStore } -func New(ds model.DataStore, artwork artwork.Artwork, streamer stream.MediaStreamer, share core.Share, archiver core.Archiver) *Router { - p := &Router{ds: ds, artwork: artwork, streamer: streamer, share: share, archiver: archiver} +func New(ds model.DataStore, artwork artwork.Artwork, streamer stream.MediaStreamer, decider stream.TranscodeDecider, share core.Share, archiver core.Archiver) *Router { + p := &Router{ds: ds, artwork: artwork, streamer: streamer, decider: decider, share: share, archiver: archiver} shareRoot := path.Join(conf.Server.BasePath, consts.URLPathPublic) p.assetsHandler = http.StripPrefix(shareRoot, http.FileServer(http.FS(ui.BuildAssets()))) p.Handler = p.routes() diff --git a/server/serve_index.go b/server/serve_index.go index a538daf1a..167197403 100644 --- a/server/serve_index.go +++ b/server/serve_index.go @@ -14,6 +14,7 @@ import ( "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/conf/mime" "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/core/quickconnect" "github.com/navidrome/navidrome/log" "github.com/navidrome/navidrome/model" "github.com/navidrome/navidrome/utils/slice" @@ -31,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) @@ -57,6 +58,8 @@ func serveIndex(ds model.DataStore, fs fs.FS, shareInfo *model.Share) http.Handl "uiSearchDebounceMs": conf.Server.UISearchDebounceMs, "uiCoverArtSize": conf.Server.UICoverArtSize, "enableCoverAnimation": conf.Server.EnableCoverAnimation, + "pidAlbum": conf.Server.PID.Album, + "pidTrack": conf.Server.PID.Track, "enableNowPlaying": conf.Server.EnableNowPlaying, "playbackReportIntervalMs": conf.Server.UIPlaybackReportInterval.Milliseconds(), "gaTrackingId": conf.Server.GATrackingID, @@ -75,6 +78,7 @@ func serveIndex(ds model.DataStore, fs fs.FS, shareInfo *model.Share) http.Handl "listenBrainzEnabled": conf.Server.ListenBrainz.Enabled, "enableExternalServices": conf.Server.EnableExternalServices, "enableReplayGain": conf.Server.EnableReplayGain, + "enableQuickConnect": quickconnect.Enabled(), "defaultDownsamplingFormat": conf.Server.DefaultDownsamplingFormat, "separator": string(os.PathSeparator), "enableInspect": conf.Server.Inspect.Enabled, diff --git a/server/serve_index_test.go b/server/serve_index_test.go index 78f3873b8..277513768 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" @@ -88,6 +89,8 @@ var _ = Describe("serveIndex", func() { Entry("uiSearchDebounceMs", func() { conf.Server.UISearchDebounceMs = 500 }, "uiSearchDebounceMs", float64(500)), Entry("uiCoverArtSize", func() { conf.Server.UICoverArtSize = 300 }, "uiCoverArtSize", float64(300)), Entry("enableCoverAnimation", func() { conf.Server.EnableCoverAnimation = true }, "enableCoverAnimation", true), + Entry("pidAlbum", func() { conf.Server.PID.Album = "folder" }, "pidAlbum", "folder"), + Entry("pidTrack", func() { conf.Server.PID.Track = "title" }, "pidTrack", "title"), Entry("enableNowPlaying", func() { conf.Server.EnableNowPlaying = true }, "enableNowPlaying", true), Entry("gaTrackingId", func() { conf.Server.GATrackingID = "UA-12345" }, "gaTrackingId", "UA-12345"), Entry("defaultDownloadableShare", func() { conf.Server.DefaultDownloadableShare = true }, "defaultDownloadableShare", true), @@ -97,6 +100,8 @@ var _ = Describe("serveIndex", func() { Entry("devUIShowConfig", func() { conf.Server.DevUIShowConfig = true }, "devUIShowConfig", true), Entry("listenBrainzEnabled", func() { conf.Server.ListenBrainz.Enabled = true }, "listenBrainzEnabled", true), Entry("enableReplayGain", func() { conf.Server.EnableReplayGain = true }, "enableReplayGain", true), + Entry("enableQuickConnect", func() { conf.Server.Jellyfin.Enabled = true; conf.Server.Jellyfin.QuickConnect = true }, "enableQuickConnect", true), + Entry("enableQuickConnect without the Jellyfin API", func() { conf.Server.Jellyfin.Enabled = false; conf.Server.Jellyfin.QuickConnect = true }, "enableQuickConnect", false), Entry("enableExternalServices", func() { conf.Server.EnableExternalServices = true }, "enableExternalServices", true), Entry("devActivityPanel", func() { conf.Server.DevActivityPanel = true }, "devActivityPanel", true), Entry("shareURL", func() { conf.Server.ShareURL = "https://share.example.com" }, "shareURL", "https://share.example.com"), @@ -339,7 +344,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/server.go b/server/server.go index b05c20cc5..7e7e2f6bd 100644 --- a/server/server.go +++ b/server/server.go @@ -17,7 +17,6 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5/middleware" - "github.com/go-chi/httprate" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/core/auth" @@ -205,11 +204,12 @@ func (s *Server) initRoutes() { func (s *Server) mountAuthenticationRoutes() chi.Router { r := s.router return r.Route(path.Join(conf.Server.BasePath, "/auth"), func(r chi.Router) { + r.Use(LimitLoginBody) if conf.Server.AuthRequestLimit > 0 { log.Info("Login rate limit set", "requestLimit", conf.Server.AuthRequestLimit, "windowLength", conf.Server.AuthWindowLength) - rateLimiter := httprate.LimitByIP(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength) + rateLimiter := ClientIPRateLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength) r.With(rateLimiter).Post("/login", login(s.ds)) } else { log.Warn("Login rate limit is disabled! Consider enabling it to be protected against brute-force attacks") 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/api.go b/server/subsonic/api.go index 029046c39..fe724741c 100644 --- a/server/subsonic/api.go +++ b/server/subsonic/api.go @@ -9,7 +9,6 @@ import ( "regexp" "strconv" - "github.com/deluan/rest" "github.com/go-chi/chi/v5" "github.com/navidrome/navidrome/conf" "github.com/navidrome/navidrome/core" @@ -108,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)) @@ -302,10 +302,9 @@ func mapToSubsonicError(err error) subError { err = newError(responses.ErrorMissingParameter, err.Error()) case errors.Is(err, req.ErrInvalidParam): err = newError(responses.ErrorGeneric, err.Error()) - case errors.Is(err, model.ErrNotFound), errors.Is(err, rest.ErrNotFound): + case errors.Is(err, model.ErrNotFound): err = newError(responses.ErrorDataNotFound, "data not found") - case errors.Is(err, model.ErrNotAuthorized), errors.Is(err, rest.ErrPermissionDenied), - errors.Is(err, model.ErrPlaylistNotEditable): // Subsonic has no code for "read-only resource" + case errors.Is(err, model.ErrNotAuthorized), errors.Is(err, model.ErrPlaylistNotEditable): // Subsonic has no code for "read-only resource" err = newError(responses.ErrorAuthorizationFail) case errors.Is(err, stream.ErrTooManyTranscodes): err = newError(responses.ErrorGeneric, "too many concurrent transcodes, please retry shortly") diff --git a/server/subsonic/auth_limiter.go b/server/subsonic/auth_limiter.go new file mode 100644 index 000000000..88df6eb50 --- /dev/null +++ b/server/subsonic/auth_limiter.go @@ -0,0 +1,134 @@ +package subsonic + +import ( + "cmp" + "context" + "hash/maphash" + "sync" + "time" + + "github.com/navidrome/navidrome/consts" +) + +// authLimiter caps failed Subsonic logins per key. Checks run at most `limit` at a time and failures +// are recorded afterwards, so a window admits up to 2*limit-1 guesses and valid requests only wait. +type authLimiter struct { + limit int + window time.Duration + seed maphash.Seed + mu sync.Mutex + keys map[uint64]*authAttempts // hashed, so attacker-chosen usernames cannot bloat memory + lastSweep time.Time +} + +type authAttempts struct { + failures int + start time.Time + slots chan struct{} + refs int +} + +// authSlot is a reserved credential check. A nil slot releases nothing, which is what a disabled +// limiter hands back. +type authSlot struct { + limiter *authLimiter + entry *authAttempts +} + +// newAuthLimiter returns nil when limit is not positive. A nil limiter allows everything. +func newAuthLimiter(limit int, window time.Duration) *authLimiter { + if limit <= 0 { + return nil + } + return &authLimiter{ + limit: limit, + window: cmp.Or(window, consts.DefaultAuthWindowLength), + seed: maphash.MakeSeed(), + keys: map[uint64]*authAttempts{}, + } +} + +// acquire reserves a credential check for key, waiting while other checks for the same key are in +// flight. It only fails when the key already reached `limit` failures in the current window. +func (l *authLimiter) acquire(ctx context.Context, key string) (*authSlot, bool) { + if l == nil { + return nil, true + } + a, ok := l.reserve(key) + if !ok { + return nil, false + } + + select { + case a.slots <- struct{}{}: + case <-ctx.Done(): + l.unref(a) + return nil, false + } + + l.mu.Lock() + blocked := a.failures >= l.limit + if blocked { + a.refs-- + } + l.mu.Unlock() + if blocked { + <-a.slots + return nil, false + } + return &authSlot{limiter: l, entry: a}, true +} + +func (l *authLimiter) reserve(key string) (*authAttempts, bool) { + now := time.Now() + l.mu.Lock() + defer l.mu.Unlock() + l.sweep(now) + + h := maphash.String(l.seed, key) + a := l.keys[h] + switch { + case a == nil: + a = &authAttempts{start: now, slots: make(chan struct{}, l.limit)} + l.keys[h] = a + case now.Sub(a.start) >= l.window: + a.failures, a.start = 0, now + } + if a.failures >= l.limit { + return nil, false + } + a.refs++ + return a, true +} + +func (l *authLimiter) unref(a *authAttempts) { + l.mu.Lock() + a.refs-- + l.mu.Unlock() +} + +func (s *authSlot) release(failed bool) { + if s == nil { + return + } + s.limiter.mu.Lock() + if failed { + s.entry.failures++ + } + s.entry.refs-- + s.limiter.mu.Unlock() + <-s.entry.slots +} + +// sweep drops idle expired keys once per window, so memory is bounded by recent attempts. +func (l *authLimiter) sweep(now time.Time) { + if now.Sub(l.lastSweep) < l.window { + return + } + for h, a := range l.keys { + if a.refs == 0 && now.Sub(a.start) >= l.window { + delete(l.keys, h) + } + } + l.lastSweep = now +} diff --git a/server/subsonic/auth_limiter_test.go b/server/subsonic/auth_limiter_test.go new file mode 100644 index 000000000..66e1a594f --- /dev/null +++ b/server/subsonic/auth_limiter_test.go @@ -0,0 +1,143 @@ +package subsonic + +import ( + "context" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("authLimiter", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = context.Background() + }) + + acquire := func(l *authLimiter, key string) (*authSlot, bool) { + GinkgoHelper() + return l.acquire(ctx, key) + } + + It("blocks a key after the configured number of failures", func() { + l := newAuthLimiter(2, time.Minute) + for range 2 { + slot, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + slot.release(true) + } + + _, ok := acquire(l, "k") + Expect(ok).To(BeFalse()) + }) + + It("never counts successful checks", func() { + l := newAuthLimiter(2, time.Minute) + for range 50 { + slot, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + slot.release(false) + } + }) + + It("keeps keys independent", func() { + l := newAuthLimiter(1, time.Minute) + slot, _ := acquire(l, "a") + slot.release(true) + _, ok := acquire(l, "a") + Expect(ok).To(BeFalse()) + + _, ok = acquire(l, "b") + Expect(ok).To(BeTrue()) + }) + + It("waits for an in-flight check instead of failing the request", func() { + l := newAuthLimiter(1, time.Minute) + held, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + + waiting := make(chan bool, 1) + go func() { + slot, ok := l.acquire(ctx, "k") + slot.release(false) + waiting <- ok + }() + Consistently(waiting, 50*time.Millisecond).ShouldNot(Receive()) + + held.release(false) + Eventually(waiting).Should(Receive(BeTrue())) + }) + + It("stops waiting when the request is canceled", func() { + l := newAuthLimiter(1, time.Minute) + held, _ := acquire(l, "k") + DeferCleanup(func() { held.release(false) }) + + canceled, cancel := context.WithCancel(context.Background()) + cancel() + _, ok := l.acquire(canceled, "k") + Expect(ok).To(BeFalse()) + }) + + It("does not let concurrent guesses overshoot the limit", func() { + l := newAuthLimiter(5, time.Minute) + hold := make(chan struct{}) + var checks atomic.Int32 + var wg sync.WaitGroup + for range 50 { + wg.Go(func() { + slot, ok := l.acquire(ctx, "k") + if !ok { + return + } + checks.Add(1) + <-hold + slot.release(true) + }) + } + + Eventually(checks.Load).Should(Equal(int32(5))) + Consistently(checks.Load, 100*time.Millisecond).Should(Equal(int32(5))) + close(hold) + wg.Wait() + + _, ok := acquire(l, "k") + Expect(ok).To(BeFalse()) + }) + + It("allows everything when the limit is disabled", func() { + l := newAuthLimiter(0, time.Minute) + for range 10 { + slot, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + slot.release(true) + } + }) +}) + +// testing/synctest's fake clock needs a *testing.T, which Ginkgo doesn't give. +func TestAuthLimiterWindow(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + g := NewWithT(t) + ctx := context.Background() + l := newAuthLimiter(1, 20*time.Second) + for _, key := range []string{"a", "b"} { + slot, ok := l.acquire(ctx, key) + g.Expect(ok).To(BeTrue()) + slot.release(true) + } + _, ok := l.acquire(ctx, "a") + g.Expect(ok).To(BeFalse()) + + time.Sleep(20 * time.Second) + slot, ok := l.acquire(ctx, "a") + g.Expect(ok).To(BeTrue()) + slot.release(false) + g.Expect(l.keys).To(HaveLen(1), "expired keys must be dropped") + }) +} diff --git a/server/subsonic/bookmarks.go b/server/subsonic/bookmarks.go index 4a7ebaa6c..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,16 @@ 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()) - err = repo.AddBookmark(id, comment, position) + repo := api.ds.MediaFile() + ok, err := repo.Exists(r.Context(), id) + if err != nil { + return nil, err + } + if !ok { + return nil, newError(responses.ErrorDataNotFound, "Song not found") + } + + err = repo.AddBookmark(r.Context(), id, comment, position) if err != nil { return nil, err } @@ -61,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 } @@ -72,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 } @@ -132,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 } @@ -143,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 } @@ -207,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 new file mode 100644 index 000000000..387ba6ab2 --- /dev/null +++ b/server/subsonic/bookmarks_test.go @@ -0,0 +1,48 @@ +package subsonic + +import ( + "context" + + "github.com/navidrome/navidrome/model" + "github.com/navidrome/navidrome/model/request" + "github.com/navidrome/navidrome/server/subsonic/responses" + "github.com/navidrome/navidrome/tests" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("Bookmarks", func() { + var router *Router + var ds *tests.MockDataStore + var mfRepo *tests.MockMediaFileRepo + var ctx context.Context + + BeforeEach(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().(*tests.MockMediaFileRepo) + mfRepo.SetData(model.MediaFiles{{ID: "visible"}}) + }) + + Describe("CreateBookmark", func() { + It("rejects an id the user cannot read", func() { + r := newGetRequest("id=hidden", "position=1").WithContext(ctx) + + _, err := router.CreateBookmark(r) + + Expect(err).To(HaveOccurred()) + Expect(mapToSubsonicError(err).code).To(Equal(responses.ErrorDataNotFound)) + Expect(mfRepo.BookmarksAdded).To(BeEmpty()) + }) + + It("accepts an id the user can read", func() { + r := newGetRequest("id=visible", "position=1").WithContext(ctx) + + _, err := router.CreateBookmark(r) + + Expect(err).ToNot(HaveOccurred()) + Expect(mfRepo.BookmarksAdded).To(ConsistOf("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_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/e2e/subsonic_artwork_test.go b/server/subsonic/e2e/subsonic_artwork_test.go index 9324ea9e3..530588759 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()) @@ -137,7 +137,7 @@ var _ = Describe("Artwork Serving", Ordered, func() { artRouter = buildArtworkRouter(artSvc) router = artRouter // so the shared doReq/doRawReq helpers hit the artwork-wired router - pubRouter = public.New(ds, artSvc, streamerSpy, core.NewShare(ds), noopArchiver{}) + pubRouter = public.New(ds, artSvc, streamerSpy, stream.NewTranscodeDecider(ds, ffm), core.NewShare(ds), noopArchiver{}) }) It("emits a bare optimistic coverArt id before the queue is drained", func() { @@ -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 079b131ff..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 @@ -106,14 +106,15 @@ 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 - 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 + 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/e2e/subsonic_multilibrary_test.go b/server/subsonic/e2e/subsonic_multilibrary_test.go index 98f87ca17..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", }) @@ -224,13 +224,31 @@ var _ = Describe("Multi-Library Support", Ordered, func() { Expect(resp.Playlist.Entry).To(HaveLen(1)) Expect(resp.Playlist.Entry[0].Id).To(Equal(lib1SongID)) }) + + It("non-admin user cannot store a song from another library through createPlaylist", func() { + resp := doReqWithUser(userLib1Only, "createPlaylist", + "name", "Restricted Playlist", "songId", lib1SongID, "songId", lib2SongID) + Expect(resp.Status).To(Equal(responses.StatusOK)) + ownID := resp.Playlist.Id + + stored := doReqWithUser(adminWithLibs, "getPlaylist", "id", ownID) + Expect(stored.Playlist.Entry).To(HaveLen(1), "the lib2 song must not be persisted") + Expect(stored.Playlist.Entry[0].Id).To(Equal(lib1SongID)) + + By("replacing the tracks of the same playlist") + resp = doReqWithUser(userLib1Only, "createPlaylist", "playlistId", ownID, "songId", lib2SongID) + Expect(resp.Status).To(Equal(responses.StatusOK)) + + stored = doReqWithUser(adminWithLibs, "getPlaylist", "id", ownID) + Expect(stored.Playlist.Entry).To(BeEmpty()) + }) }) Describe("Cross-library shares", 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_radio_test.go b/server/subsonic/e2e/subsonic_radio_test.go index cd778fa79..d51cf6b14 100644 --- a/server/subsonic/e2e/subsonic_radio_test.go +++ b/server/subsonic/e2e/subsonic_radio_test.go @@ -130,4 +130,12 @@ var _ = Describe("Internet Radio Endpoints", Ordered, func() { Expect(resp.InternetRadioStations).ToNot(BeNil()) Expect(resp.InternetRadioStations.Radios).To(BeEmpty()) }) + + It("deleteInternetRadioStation returns not found for a missing station", func() { + resp := doReq("deleteInternetRadioStation", "id", radioID) + + Expect(resp.Status).To(Equal(responses.StatusFailed)) + Expect(resp.Error).ToNot(BeNil()) + Expect(resp.Error.Code).To(Equal(responses.ErrorDataNotFound)) + }) }) 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/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/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 cfbff3ecb..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) @@ -181,6 +181,9 @@ func (api *Router) Scrobble(r *http.Request) (*responses.Subsonic, error) { log.Error(ctx, "Error registering scrobbles", "ids", ids, "times", times, err) } } else { + if len(ids) > 1 { + log.Warn(ctx, "Multiple ids sent to a nowPlaying notification, only the first one will be used", "ids", ids) + } err := api.scrobblerNowPlaying(ctx, ids[0], position) if err != nil { log.Error(ctx, "Error setting NowPlaying", "id", ids[0], err) @@ -207,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 6dfa2263f..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) @@ -98,17 +99,22 @@ func checkRequiredParameters(next http.Handler) http.Handler { } func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { + limiter := newAuthLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength) return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { 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(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 @@ -118,29 +124,51 @@ 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") salt, _ := p.String("s") jwt, _ := p.String("jwt") - usr, err = ds.User(ctx).FindByUsernameWithPassword(username) - if errors.Is(err, context.Canceled) { - log.Debug(ctx, "API: Request canceled when authenticating", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) + // Blocked attempts get the same response as a wrong password, so they reveal nothing + limitKey := server.ClientIP(r) + "\x00" + strings.ToLower(username) + slot, allowed := limiter.acquire(ctx, limitKey) + if !allowed { + if ctx.Err() != nil { + return + } + log.Warn(ctx, "API: Too many failed login attempts", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr) + sendError(w, r, newError(responses.ErrorAuthenticationFail)) return } + + 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) switch { - case errors.Is(err, model.ErrNotFound): + case errors.Is(err, context.Canceled): + log.Debug(ctx, "API: Request canceled when authenticating", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) + return + case invalidLogin: log.Warn(ctx, "API: Invalid login", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) case err != nil: log.Error(ctx, "API: Error authenticating username", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) - default: - err = validateCredentials(usr, pass, token, salt, jwt) - if err != nil { - log.Warn(ctx, "API: Invalid login", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) - } } } @@ -150,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()) @@ -182,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 @@ -205,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 { @@ -218,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 cb34b92e7..b3ad972b7 100644 --- a/server/subsonic/middlewares_test.go +++ b/server/subsonic/middlewares_test.go @@ -3,11 +3,14 @@ package subsonic import ( "context" "crypto/md5" + "encoding/hex" "errors" "fmt" "net/http" "net/http/httptest" "strings" + "sync" + "sync/atomic" "time" "github.com/navidrome/navidrome/conf" @@ -40,11 +43,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{} @@ -115,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) @@ -145,8 +158,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", }) @@ -306,6 +319,295 @@ var _ = Describe("Middlewares", func() { Expect(next.called).To(BeFalse()) }) }) + + 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 + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.AuthRequestLimit = 3 + conf.Server.AuthWindowLength = time.Minute + cp = authenticate(ds)(next) + }) + + serve := func(r *http.Request) *httptest.ResponseRecorder { + next.called = false + rec := httptest.NewRecorder() + cp.ServeHTTP(rec, r) + return rec + } + failTimes := func(n int, params ...string) { + for range n { + Expect(serve(newGetRequest(params...)).Body.String()).To(ContainSubstring(`code="40"`)) + } + } + + It("rejects the correct password exactly like a wrong one", func() { + failTimes(3, "u=admin", "p=WRONG") + + rec := serve(newGetRequest("u=admin", "p=wordpass")) + + Expect(next.called).To(BeFalse()) + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(rec.Body.String()).To(ContainSubstring(`code="40"`)) + Expect(rec.Header().Get("Retry-After")).To(BeEmpty()) + }) + + It("counts attempts against unknown usernames", func() { + failTimes(3, "u=newuser", "p=secret") + _ = ds.User().Put(ctx, &model.User{UserName: "newuser", NewPassword: "secret"}) + + serve(newGetRequest("u=newuser", "p=secret")) + Expect(next.called).To(BeFalse()) + }) + + It("treats usernames case-insensitively", func() { + failTimes(3, "u=ADMIN", "p=WRONG") + + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeFalse()) + }) + + It("does not count successful logins", func() { + for range 10 { + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeTrue()) + } + }) + + It("does not count server errors", func() { + userRepo := ds.User().(*tests.MockedUserRepo) + userRepo.Error = errors.New("db down") + failTimes(5, "u=admin", "p=wordpass") + userRepo.Error = nil + + serve(newGetRequest("u=admin", "p=wordpass")) + 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") + + serve(newGetRequest("u=other", "p=otherpass")) + Expect(next.called).To(BeTrue()) + }) + + It("does not block the same username from another IP", func() { + failTimes(3, "u=admin", "p=WRONG") + + r := newGetRequest("u=admin", "p=wordpass") + r.RemoteAddr = "198.51.100.7:1234" + serve(r) + Expect(next.called).To(BeTrue()) + }) + + It("does not limit reverse proxy authentication", func() { + conf.Server.ExtAuth.TrustedSources = "192.168.1.1/24" + conf.Server.ExtAuth.UserHeader = "Remote-User" + failTimes(3, "u=admin", "p=WRONG") + + r := newGetRequest() + r.Header.Add("Remote-User", "admin") + r = r.WithContext(request.WithReverseProxyIp(r.Context(), "192.168.1.1")) + serve(r) + 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) + failTimes(10, "u=admin", "p=WRONG") + + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeTrue()) + }) + }) + + When("valid requests overlap", func() { + var gate *gatedUserRepo + var gatedDS model.DataStore + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.AuthRequestLimit = 5 + conf.Server.AuthWindowLength = time.Minute + gate = &gatedUserRepo{ + UserRepository: ds.User(), + entered: make(chan struct{}, 64), + proceed: make(chan struct{}), + } + gatedDS = &gatedDataStore{DataStore: ds, users: gate} + }) + + It("lets every valid request through while checks are in flight", func() { + const burst = 6 + cp := authenticate(gatedDS)(&countingHandler{}) + var passed atomic.Int32 + var wg sync.WaitGroup + for range burst { + wg.Go(func() { + rec := httptest.NewRecorder() + cp.ServeHTTP(rec, newGetRequest("u=admin", "p=wordpass")) + if !strings.Contains(rec.Body.String(), `code="40"`) { + passed.Add(1) + } + }) + } + for range conf.Server.AuthRequestLimit { + Eventually(gate.entered).Should(Receive()) + } + close(gate.proceed) + wg.Wait() + + Expect(passed.Load()).To(Equal(int32(burst))) + }) + + It("caps concurrent credential checks for wrong passwords", func() { + cp := authenticate(gatedDS)(&countingHandler{}) + var wg sync.WaitGroup + for i := range 100 { + wg.Go(func() { + cp.ServeHTTP(httptest.NewRecorder(), newGetRequest("u=admin", fmt.Sprintf("p=wrong%d", i))) + }) + } + + limit := int32(conf.Server.AuthRequestLimit) + Eventually(gate.lookups.Load).Should(Equal(limit)) + Consistently(gate.lookups.Load, 100*time.Millisecond).Should(Equal(limit)) + close(gate.proceed) + wg.Wait() + }) + }) }) Describe("AdminOnly", func() { @@ -368,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{ @@ -421,14 +741,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) } @@ -548,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) { @@ -560,3 +886,28 @@ func (mp *mockPlayers) Register(ctx context.Context, id, client, typ, ip string) } return &model.Player{ID: id}, mp.transcoding, nil } + +type gatedDataStore struct { + model.DataStore + users model.UserRepository +} + +func (g *gatedDataStore) User() model.UserRepository { return g.users } + +type gatedUserRepo struct { + model.UserRepository + entered chan struct{} + proceed chan struct{} + lookups atomic.Int32 +} + +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(ctx, username) +} + +type countingHandler struct{ calls atomic.Int32 } + +func (c *countingHandler) ServeHTTP(http.ResponseWriter, *http.Request) { c.calls.Add(1) } 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/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/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/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 28b4585f0..fd26ccc4d 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) { @@ -28,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 } @@ -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/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/server/subsonic/transcode.go b/server/subsonic/transcode.go index 9eb2af160..7a011a616 100644 --- a/server/subsonic/transcode.go +++ b/server/subsonic/transcode.go @@ -280,12 +280,19 @@ func (api *Router) GetTranscodeDecision(w http.ResponseWriter, r *http.Request) return stream.IsAACCodec(p.Container) }) + player, hasPlayer := request.PlayerFrom(ctx) + // Honor the player's forced transcoding format, falling back to normal // negotiation when the client can't play it (issue #5583). + maxBitRate := 0 if trc, ok := request.TranscodingFrom(ctx); ok && trc.TargetFormat != "" { - if !clientInfo.ForceFormat(trc.TargetFormat) { + if clientInfo.ForceFormat(trc.TargetFormat) { + // DirectPlayProfile carries no bitrate, so this ceiling is the only + // thing keeping an over-bitrate source out of direct play. + maxBitRate = trc.DefaultBitRate + } else { clientName := clientInfo.Name - if player, ok := request.PlayerFrom(ctx); ok && player.Client != "" { + if hasPlayer && player.Client != "" { clientName = player.Client } log.Debug(ctx, "Player forced format not supported by client; falling back to negotiation", @@ -293,17 +300,17 @@ func (api *Router) GetTranscodeDecision(w http.ResponseWriter, r *http.Request) } } - // Apply the player's MaxBitRate as a ceiling on the client's declared - // limits (issue #5583). Both fields are capped because the client sends - // them independently here; capping only MaxAudioBitrate would let an - // independent MaxTranscodingAudioBitrate slip through computeBitrate. - if player, ok := request.PlayerFrom(ctx); ok && clientInfo.CapBitrate(player.MaxBitRate) { - log.Debug(ctx, "Applied player MaxBitRate cap to transcode decision", - "playerMaxBitRate", player.MaxBitRate, "client", clientInfo.Name) + // The player's own MaxBitRate outranks the forced-format default (issue #5583). + if hasPlayer && player.MaxBitRate > 0 { + maxBitRate = player.MaxBitRate + } + if clientInfo.CapBitrate(maxBitRate) { + log.Debug(ctx, "Applied bitrate ceiling to transcode decision", + "maxBitRate", maxBitRate, "client", clientInfo.Name) } // 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) @@ -392,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/server/subsonic/transcode_test.go b/server/subsonic/transcode_test.go index 8d5cbb974..1f2fdfcef 100644 --- a/server/subsonic/transcode_test.go +++ b/server/subsonic/transcode_test.go @@ -369,7 +369,7 @@ var _ = Describe("Transcode endpoints", func() { mockTD.token = "token" }) - It("forces a supported format and clears direct play", func() { + It("forces a supported format and narrows direct play to it", func() { body := `{"directPlayProfiles":[{"containers":["flac"],"audioCodecs":["flac"],"protocols":["http"]}], "transcodingProfiles":[{"container":"ogg","audioCodec":"opus","protocol":"http"}, {"container":"mp3","audioCodec":"mp3","protocol":"http"}]}` @@ -380,7 +380,11 @@ var _ = Describe("Transcode endpoints", func() { Expect(err).ToNot(HaveOccurred()) Expect(mockTD.capturedClient.TranscodingProfiles).To(HaveLen(1)) Expect(mockTD.capturedClient.TranscodingProfiles[0].AudioCodec).To(Equal("opus")) - Expect(mockTD.capturedClient.DirectPlayProfiles).To(BeEmpty()) + Expect(mockTD.capturedClient.DirectPlayProfiles).To(ConsistOf(stream.DirectPlayProfile{ + Containers: []string{"ogg"}, + AudioCodecs: []string{"opus"}, + Protocols: []string{"http"}, + })) }) It("falls back to negotiation when the forced format is unsupported", func() { @@ -416,6 +420,43 @@ var _ = Describe("Transcode endpoints", func() { Expect(mockTD.capturedClient.MaxAudioBitrate).To(Equal(128)) Expect(mockTD.capturedClient.MaxTranscodingAudioBitrate).To(Equal(128)) }) + + withForcedBitRate := func(r *http.Request, format string, defaultBitRate, playerMaxBitRate int) *http.Request { + ctx := request.WithTranscoding(r.Context(), model.Transcoding{TargetFormat: format, DefaultBitRate: defaultBitRate}) + ctx = request.WithPlayer(ctx, model.Player{Client: "NavidromeUI", MaxBitRate: playerMaxBitRate}) + return r.WithContext(ctx) + } + + It("applies the transcoding default bitrate when the player sets no maxBitRate", func() { + body := `{"transcodingProfiles":[{"container":"mp3","audioCodec":"mp3","protocol":"http"}]}` + r := withForcedBitRate(newJSONPostRequest("mediaId=song-1&mediaType=song", body), "mp3", 192, 0) + + _, err := router.GetTranscodeDecision(w, r) + + Expect(err).ToNot(HaveOccurred()) + Expect(mockTD.capturedClient.MaxAudioBitrate).To(Equal(192)) + Expect(mockTD.capturedClient.MaxTranscodingAudioBitrate).To(Equal(192)) + }) + + It("prefers the player maxBitRate over the transcoding default bitrate", func() { + body := `{"transcodingProfiles":[{"container":"mp3","audioCodec":"mp3","protocol":"http"}]}` + r := withForcedBitRate(newJSONPostRequest("mediaId=song-1&mediaType=song", body), "mp3", 192, 320) + + _, err := router.GetTranscodeDecision(w, r) + + Expect(err).ToNot(HaveOccurred()) + Expect(mockTD.capturedClient.MaxAudioBitrate).To(Equal(320)) + }) + + It("ignores the transcoding default bitrate when the forced format is unsupported", func() { + body := `{"transcodingProfiles":[{"container":"mp3","audioCodec":"mp3","protocol":"http"}]}` + r := withForcedBitRate(newJSONPostRequest("mediaId=song-1&mediaType=song", body), "opus", 192, 0) + + _, err := router.GetTranscodeDecision(w, r) + + Expect(err).ToNot(HaveOccurred()) + Expect(mockTD.capturedClient.MaxAudioBitrate).To(BeZero()) + }) }) }) @@ -584,6 +625,10 @@ func (m *mockTranscodeDecision) ResolveRequest(_ context.Context, _ *model.Media return stream.Request{Format: "raw"} } +func (m *mockTranscodeDecision) ResolveClientRequest(context.Context, *model.MediaFile, *stream.ClientInfo, int) stream.Request { + return stream.Request{Format: "raw"} +} + func (m *mockTranscodeDecision) CreateTranscodeParams(_ *stream.TranscodeDecision) (string, error) { return m.token, m.tokenErr } 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...)) +} diff --git a/tests/fixtures/lastfm.artist.page.challenge.html b/tests/fixtures/lastfm.artist.page.challenge.html new file mode 100644 index 000000000..f431891bd --- /dev/null +++ b/tests/fixtures/lastfm.artist.page.challenge.html @@ -0,0 +1,88 @@ + + + + + + + + Client Challenge + + + + + + + + diff --git a/tests/harness/harness.go b/tests/harness/harness.go index 5949c4fae..67f3150a2 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 } @@ -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/init_tests.go b/tests/init_tests.go index 902cf196d..eee4428a8 100644 --- a/tests/init_tests.go +++ b/tests/init_tests.go @@ -8,6 +8,7 @@ import ( "testing" "github.com/navidrome/navidrome/conf" + _ "github.com/navidrome/navidrome/conf/mime" // registers mime_types.yaml, so tests see the same image types as the server "github.com/navidrome/navidrome/log" ) 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 32f56a4f0..a5e4126cf 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 }{} + db.MockedPlayer = CreateMockPlayerRepo() 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 @@ -329,35 +329,8 @@ func (db *MockDataStore) WithTxImmediate(block func(tx model.DataStore) error, l return block(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) WithTxRetry(ctx context.Context, block func(ctx context.Context, tx model.DataStore) error, label ...string) error { + return block(ctx, db) } func (db *MockDataStore) GC(context.Context, ...int) error { 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 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 1a16a7e0b..6de2b9265 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,59 @@ 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) SetScannedPID(_ context.Context, id int, pid model.PIDConfig) error { + if m.Err != nil { + return m.Err + } + if lib, ok := m.Data[id]; ok { + lib.ScannedPIDAlbum, lib.ScannedPIDTrack = pid.Album, pid.Track + m.Data[id] = lib + } + return nil +} + +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,35 +177,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) - mf, err := m.Get(idInt) - if errors.Is(err, model.ErrNotFound) { - return nil, rest.ErrNotFound - } - return mf, err + 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 } @@ -218,8 +216,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 } @@ -311,4 +309,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 7a1a8f926..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" @@ -24,22 +25,30 @@ type MockMediaFileRepo struct { model.MediaFileRepository Data map[string]*model.MediaFile Err bool - // Add fields and methods for controlling CountAll and DeleteAllMissing in tests - CountAllValue int64 + // Add fields and methods for controlling CountAll and DeleteAllMissing in tests. + // A nil CountAllValue is unset, and CountAll falls back to counting rows in Data. + CountAllValue *int64 CountAllOptions model.QueryOptions DeleteAllMissingValue int64 - Options model.QueryOptions + // ReassignReferencesCalls records prevID -> newID + ReassignReferencesCalls map[string]string + Options model.QueryOptions // Add fields for cross-library move detection tests FindRecentFilesByMBZTrackIDFunc func(missing model.MediaFile, since time.Time) (model.MediaFiles, error) FindRecentFilesByPropertiesFunc func(missing model.MediaFile, since time.Time) (model.MediaFiles, error) MatchesCriteriaValue bool MatchesCriteriaErr error + BookmarksAdded []string } func (m *MockMediaFileRepo) SetError(err bool) { m.Err = err } +func (m *MockMediaFileRepo) SetCountAll(count int64) { + m.CountAllValue = &count +} + func (m *MockMediaFileRepo) SetData(mfs model.MediaFiles) { m.Data = make(map[string]*model.MediaFile) for i, mf := range mfs { @@ -47,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") } @@ -55,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") } @@ -69,7 +78,15 @@ func (m *MockMediaFileRepo) Get(id string) (*model.MediaFile, error) { return nil, model.ErrNotFound } -func (m *MockMediaFileRepo) GetWithParticipants(id string) (*model.MediaFile, error) { +func (m *MockMediaFileRepo) AddBookmark(_ context.Context, id, _ string, _ int64) error { + if m.Err { + return errors.New("error") + } + m.BookmarksAdded = append(m.BookmarksAdded, id) + return nil +} + +func (m *MockMediaFileRepo) GetWithParticipants(_ context.Context, id string) (*model.MediaFile, error) { if m.Err { return nil, errors.New("error") } @@ -79,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] } @@ -101,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 } @@ -112,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 } @@ -126,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") } @@ -141,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") } @@ -152,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") } @@ -163,7 +180,18 @@ func (m *MockMediaFileRepo) Delete(id string) error { return nil } -func (m *MockMediaFileRepo) IncPlayCount(id string, timestamp time.Time) error { +func (m *MockMediaFileRepo) ReassignReferences(_ context.Context, prevID, newID string) error { + if m.Err { + return errors.New("error") + } + if m.ReassignReferencesCalls == nil { + m.ReassignReferencesCalls = make(map[string]string) + } + m.ReassignReferencesCalls[prevID] = newID + return nil +} + +func (m *MockMediaFileRepo) IncPlayCount(_ context.Context, id string, timestamp time.Time) error { if m.Err { return errors.New("error") } @@ -175,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") } @@ -187,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") } @@ -213,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") } @@ -247,20 +275,20 @@ 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") } - if m.CountAllValue != 0 { - if len(opts) > 0 { - m.CountAllOptions = opts[0] - } - return m.CountAllValue, nil + if len(opts) > 0 { + m.CountAllOptions = opts[0] + } + if m.CountAllValue != nil { + return *m.CountAllValue, nil } 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") } @@ -278,32 +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) { - mf, err := m.Get(id) - if errors.Is(err, model.ErrNotFound) { - return nil, rest.ErrNotFound - } - return mf, err +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] } @@ -311,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") } @@ -337,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") } @@ -363,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 } @@ -371,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_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/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 83751bb28..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" ) @@ -13,6 +15,7 @@ type MockPlaylistTrackRepo struct { DeletedIds []string Reordered bool AddCount int + InsertPos int Err error AlbumIDs []string // stubbed result for GetAlbumIDs, ignoring options } @@ -39,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 } @@ -67,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 @@ -89,33 +92,38 @@ func (m *MockPlaylistTrackRepo) Add(ids []string) (int, error) { return m.AddCount, nil } -func (m *MockPlaylistTrackRepo) AddAlbums(_ []string) (int, error) { +func (m *MockPlaylistTrackRepo) Insert(ctx context.Context, ids []string, pos int) (int, error) { + m.InsertPos = pos + return m.Add(ctx, ids) +} + +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 5d22c26aa..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,15 +69,11 @@ func (m *MockPluginRepo) Get(id string) (*model.Plugin, error) { return nil, model.ErrNotFound } -func (m *MockPluginRepo) Read(id string) (any, error) { - p, err := m.Get(id) - if errors.Is(err, model.ErrNotFound) { - return nil, rest.ErrNotFound - } - return p, err +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 } @@ -109,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 } @@ -127,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] } @@ -140,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] } @@ -153,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} } 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' ? ( { const isXsmall = useMediaQuery((theme) => theme.breakpoints.down('xs')) const translate = useTranslate() const record = useRecordContext(props) + const locale = useDateLocale() // Create an array of detail elements let details = [] @@ -161,12 +167,13 @@ export const Details = (props) => { // Calculate date related fields const yearRange = formatRange(record, 'year') - const date = record.date ? formatFullDate(record.date) : yearRange + const date = record.date ? formatFullDate(record.date, locale) : yearRange const originalDate = record.originalDate - ? formatFullDate(record.originalDate) + ? formatFullDate(record.originalDate, locale) : formatRange(record, 'originalYear') - const releaseDate = record?.releaseDate && formatFullDate(record.releaseDate) + const releaseDate = + record?.releaseDate && formatFullDate(record.releaseDate, locale) const dateToUse = originalDate || date const isOriginalDate = originalDate && dateToUse !== date @@ -247,7 +254,12 @@ const AlbumDetails = (props) => { return (
-
+
({ + 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( + + + + + , + ) + + 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/) + }) +}) 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) => { /> + {config.enableFavourites && ( { removeItems() - if (config.extAuthLogoutURL) { + // Only proxy-authenticated sessions go to the IdP; others (e.g. direct LAN access) get the login form + if (config.extAuthLogoutURL && config.auth) { window.location.href = config.extAuthLogoutURL return Promise.resolve(false) } diff --git a/ui/src/authProvider.test.js b/ui/src/authProvider.test.js new file mode 100644 index 000000000..b16a39815 --- /dev/null +++ b/ui/src/authProvider.test.js @@ -0,0 +1,47 @@ +import { describe, it, expect, beforeEach, afterEach, vi } from 'vitest' +import config from './config' +import authProvider from './authProvider' + +vi.mock('./config', () => ({ default: {} })) + +describe('authProvider.logout', () => { + const logoutURL = 'https://auth.example.com/signout' + + beforeEach(() => { + vi.stubGlobal('location', { href: '' }) + localStorage.setItem('is-authenticated', 'true') + localStorage.setItem('token', 'abc') + }) + + afterEach(() => { + vi.unstubAllGlobals() + localStorage.clear() + delete config.extAuthLogoutURL + delete config.auth + }) + + it('clears the stored session', async () => { + await authProvider.logout() + expect(localStorage.getItem('is-authenticated')).toBeNull() + expect(localStorage.getItem('token')).toBeNull() + }) + + it('does not redirect when no logout URL is configured', async () => { + config.auth = { id: '1' } + await expect(authProvider.logout()).resolves.toBeUndefined() + expect(window.location.href).toBe('') + }) + + it('redirects to the logout URL when the page was authenticated by the proxy', async () => { + config.extAuthLogoutURL = logoutURL + config.auth = { id: '1' } + await expect(authProvider.logout()).resolves.toBe(false) + expect(window.location.href).toBe(logoutURL) + }) + + it('does not redirect when the page was not authenticated by the proxy', async () => { + config.extAuthLogoutURL = logoutURL + await expect(authProvider.logout()).resolves.toBeUndefined() + expect(window.location.href).toBe('') + }) +}) diff --git a/ui/src/common/DateField.jsx b/ui/src/common/DateField.jsx index dce24a2b9..dac9cff08 100644 --- a/ui/src/common/DateField.jsx +++ b/ui/src/common/DateField.jsx @@ -1,12 +1,14 @@ import React from 'react' import { isDateSet } from '../utils/validations' import { DateField as RADateField } from 'react-admin' +import { useDateLocale } from '../i18n/useDateLocale' export const DateField = (props) => { const { record, source } = props + const locale = useDateLocale() const value = record?.[source] if (!isDateSet(value)) return null - return + return } DateField.defaultProps = { diff --git a/ui/src/common/DateField.test.jsx b/ui/src/common/DateField.test.jsx new file mode 100644 index 000000000..47d56756b --- /dev/null +++ b/ui/src/common/DateField.test.jsx @@ -0,0 +1,32 @@ +import React from 'react' +import { render, screen } from '@testing-library/react' +import { describe, it, expect, beforeEach, vi } from 'vitest' +import { DateField } from './DateField' + +vi.mock('react-admin', async (importOriginal) => ({ + ...(await importOriginal()), + useLocale: vi.fn(), +})) + +describe('', () => { + const record = { id: '1', updatedAt: '2026-09-17T14:30:00Z' } + + beforeEach(async () => { + vi.clearAllMocks() + vi.spyOn(navigator, 'languages', 'get').mockReturnValue([]) + const { useLocale } = await import('react-admin') + vi.mocked(useLocale).mockReturnValue('de') + }) + + it('formats the date using the selected language', () => { + render() + expect(screen.getByText('17.9.2026')).toBeInTheDocument() + }) + + it('renders nothing when the date is not set', () => { + const { container } = render( + , + ) + expect(container).toBeEmptyDOMElement() + }) +}) diff --git a/ui/src/common/LoveButton.jsx b/ui/src/common/LoveButton.jsx index 492ba95b3..128b375c2 100644 --- a/ui/src/common/LoveButton.jsx +++ b/ui/src/common/LoveButton.jsx @@ -6,6 +6,8 @@ import IconButton from '@material-ui/core/IconButton' import { makeStyles } from '@material-ui/core/styles' import clsx from 'clsx' import { useToggleLove } from './useToggleLove' +import { useDateLocale } from '../i18n/useDateLocale' +import { formatDateTime } from '../utils/formatters' import { useRecordContext } from 'react-admin' import config from '../config' import { isDateSet } from '../utils/validations' @@ -40,6 +42,7 @@ export const LoveButton = ({ const record = useRecordContext({ record: recordProp }) || {} const classes = useStyles({ color, visible, loved: record.starred }) const [toggleLove, loading] = useToggleLove(resource, record) + const locale = useDateLocale() const handleToggleLove = useCallback( (e) => { @@ -61,7 +64,7 @@ export const LoveButton = ({ className={clsx(classes.love, className)} title={ isDateSet(record.starredAt) - ? new Date(record.starredAt).toLocaleString() + ? formatDateTime(record.starredAt, locale) : undefined } {...rest} diff --git a/ui/src/common/RatingField.jsx b/ui/src/common/RatingField.jsx index f92b0d948..f892b475a 100644 --- a/ui/src/common/RatingField.jsx +++ b/ui/src/common/RatingField.jsx @@ -6,6 +6,8 @@ import { isDateSet } from '../utils/validations' import StarBorderIcon from '@material-ui/icons/StarBorder' import clsx from 'clsx' import { useRating } from './useRating' +import { useDateLocale } from '../i18n/useDateLocale' +import { formatDateTime } from '../utils/formatters' import { useRecordContext } from 'react-admin' const useStyles = makeStyles({ @@ -32,6 +34,7 @@ export const RatingField = ({ const record = useRecordContext(rest) || {} const [rate, rating] = useRating(resource, record) const classes = useStyles({ color, visible }) + const locale = useDateLocale() const stopPropagation = (e) => { e.stopPropagation() @@ -50,7 +53,7 @@ export const RatingField = ({ onClick={(e) => stopPropagation(e)} title={ isDateSet(record.ratedAt) - ? new Date(record.ratedAt).toLocaleString() + ? formatDateTime(record.ratedAt, locale) : undefined } > 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 ( + } + 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 +} + +export const ReadOnlyNumberField = (props) => { + const locale = useDateLocale() + return ( + formatNumber(v, locale)} {...props} /> + ) +} + +export const ReadOnlySizeField = (props) => ( + +) + +export const ReadOnlyDurationField = (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() + + describe('', () => { + 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('', () => { + 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('', () => { + 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('', () => { + it('formats bytes as a human-readable size', () => { + renderField(ReadOnlySizeField, { source: 'size' }) + expect(screen.getByRole('textbox')).toHaveValue('1.46 MB') + }) + }) + + describe('', () => { + 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/config.js b/ui/src/config.js index 39f0cd467..62b3cb822 100644 --- a/ui/src/config.js +++ b/ui/src/config.js @@ -32,12 +32,15 @@ const defaultConfig = { listenBrainzEnabled: true, enableExternalServices: true, enableCoverAnimation: true, + pidAlbum: 'musicbrainz_albumid|albumartistid,album,albumversion,releasedate', // See consts.DefaultAlbumPID + pidTrack: 'musicbrainz_trackid|albumid,discnumber,tracknumber,title', // See consts.DefaultTrackPID enableNowPlaying: true, playbackReportIntervalMs: 60000, devShowArtistPage: true, devUIShowConfig: true, devNewEventStream: false, enableReplayGain: true, + enableQuickConnect: false, defaultDownsamplingFormat: 'opus', publicBaseUrl: '/share', separator: '/', diff --git a/ui/src/consts.js b/ui/src/consts.js index 472cd4940..46c76e194 100644 --- a/ui/src/consts.js +++ b/ui/src/consts.js @@ -31,3 +31,9 @@ export const DEFAULT_SHARE_BITRATE = 128 export const BITRATE_CHOICES = [ 32, 48, 64, 80, 96, 112, 128, 160, 192, 256, 320, ].map((b) => ({ id: b, name: b.toString() })) + +// 0 is a valid stored value ("no default bit rate") that BITRATE_CHOICES cannot express. +export const TRANSCODING_BITRATE_CHOICES = [ + { id: 0, name: 'resources.transcoding.choices.noDefaultBitRate' }, + ...BITRATE_CHOICES, +] diff --git a/ui/src/dataProvider/wrapperDataProvider.js b/ui/src/dataProvider/wrapperDataProvider.js index e79beb787..7de20bcce 100644 --- a/ui/src/dataProvider/wrapperDataProvider.js +++ b/ui/src/dataProvider/wrapperDataProvider.js @@ -148,6 +148,12 @@ const updateUser = async (params) => { 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) => { @@ -218,6 +227,15 @@ const wrapperDataProvider = { ({ json }) => ({ data: json }), ) }, + lookupQuickConnect: (code) => + httpClient( + `${REST_URL}/quickconnect?code=${encodeURIComponent(code)}`, + ).then(({ json }) => ({ data: json })), + authorizeQuickConnect: (code) => + httpClient(`${REST_URL}/quickconnect/authorize`, { + method: 'POST', + body: JSON.stringify({ code }), + }).then(({ json }) => ({ data: json })), inspect: (songId) => { return httpClient(`${REST_URL}/inspect?id=${songId}`).then(({ json }) => ({ data: json, 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/dialogs/QuickConnectDialog.jsx b/ui/src/dialogs/QuickConnectDialog.jsx new file mode 100644 index 000000000..cde7e9eaf --- /dev/null +++ b/ui/src/dialogs/QuickConnectDialog.jsx @@ -0,0 +1,128 @@ +import { useState } from 'react' +import PropTypes from 'prop-types' +import { useDataProvider, useNotify, useTranslate } from 'react-admin' +import { + Button, + Dialog, + DialogActions, + DialogContent, + DialogContentText, + DialogTitle, + TextField, +} from '@material-ui/core' + +export const QuickConnectDialog = ({ open, onClose }) => { + const translate = useTranslate() + const notify = useNotify() + const dataProvider = useDataProvider() + const [code, setCode] = useState('') + const [pending, setPending] = useState(null) + const [loading, setLoading] = useState(false) + + const handleClose = () => { + setCode('') + setPending(null) + onClose() + } + + const handleError = (error) => { + setPending(null) + const invalid = error?.status === 404 || error?.status === 409 + notify( + invalid ? 'message.quickConnectInvalidCode' : 'message.quickConnectError', + 'warning', + ) + } + + const run = (request, onSuccess) => { + setLoading(true) + request + .then(({ data }) => onSuccess(data)) + .catch(handleError) + .finally(() => setLoading(false)) + } + + const lookup = (event) => { + event.preventDefault() + run(dataProvider.lookupQuickConnect(code), setPending) + } + + const approve = () => + run(dataProvider.authorizeQuickConnect(code), (data) => { + notify('message.quickConnectApproved', 'success', { + app: data.appName, + device: data.deviceName, + }) + handleClose() + }) + + return ( + + + {translate('menu.quickConnect.name')} + + {pending ? ( + <> + + + {translate('menu.quickConnect.confirm', { + app: pending.appName, + version: pending.appVersion, + device: pending.deviceName, + })} + + + + + + + + ) : ( +
+ + + {translate('menu.quickConnect.help')} + + setCode(event.target.value)} + inputProps={{ inputMode: 'numeric', autoComplete: 'off' }} + /> + + + + + + + )} +
+ ) +} + +QuickConnectDialog.propTypes = { + open: PropTypes.bool.isRequired, + onClose: PropTypes.func.isRequired, +} diff --git a/ui/src/dialogs/QuickConnectDialog.test.jsx b/ui/src/dialogs/QuickConnectDialog.test.jsx new file mode 100644 index 000000000..5dda1cbdf --- /dev/null +++ b/ui/src/dialogs/QuickConnectDialog.test.jsx @@ -0,0 +1,103 @@ +import * as React from 'react' +import { TestContext } from 'ra-test' +import { DataProviderContext } from 'react-admin' +import { + cleanup, + fireEvent, + render, + screen, + waitFor, +} from '@testing-library/react' +import { describe, afterEach, it, expect, vi } from 'vitest' +import { QuickConnectDialog } from './QuickConnectDialog' + +const finamp = { appName: 'Finamp', appVersion: '1.0.0', deviceName: 'Pixel 7' } + +const renderDialog = (dataProvider, onClose = vi.fn()) => { + render( + + + + + , + ) + return onClose +} + +const enterCode = (code) => { + fireEvent.change(screen.getByRole('textbox'), { target: { value: code } }) + fireEvent.click(screen.getByText('menu.quickConnect.continue')) +} + +describe('QuickConnectDialog', () => { + afterEach(cleanup) + + it('shows the device before approving the code', async () => { + const dataProvider = { + lookupQuickConnect: vi.fn().mockResolvedValue({ data: finamp }), + authorizeQuickConnect: vi.fn().mockResolvedValue({ data: finamp }), + } + const onClose = renderDialog(dataProvider) + + enterCode('123 456') + await screen.findByText('menu.quickConnect.approve') + expect(dataProvider.lookupQuickConnect.mock.calls[0][0]).toBe('123 456') + expect(dataProvider.authorizeQuickConnect).not.toHaveBeenCalled() + + fireEvent.click(screen.getByText('menu.quickConnect.approve')) + await waitFor(() => expect(onClose).toHaveBeenCalled()) + expect(dataProvider.authorizeQuickConnect.mock.calls[0][0]).toBe('123 456') + }) + + it('goes back to the code input', async () => { + const dataProvider = { + lookupQuickConnect: vi.fn().mockResolvedValue({ data: finamp }), + authorizeQuickConnect: vi.fn(), + } + renderDialog(dataProvider) + + enterCode('123456') + await screen.findByText('menu.quickConnect.approve') + fireEvent.click(screen.getByText('ra.action.back')) + + expect(screen.getByRole('textbox')).toHaveValue('123456') + expect(dataProvider.authorizeQuickConnect).not.toHaveBeenCalled() + }) + + it('stays on the code input when the code is unknown', async () => { + const dataProvider = { + lookupQuickConnect: vi.fn().mockRejectedValue({ status: 404 }), + authorizeQuickConnect: vi.fn(), + } + const onClose = renderDialog(dataProvider) + + enterCode('000000') + await waitFor(() => + expect(dataProvider.lookupQuickConnect).toHaveBeenCalled(), + ) + expect(screen.getByRole('textbox')).toBeInTheDocument() + expect(screen.queryByText('menu.quickConnect.approve')).toBeNull() + expect(onClose).not.toHaveBeenCalled() + }) + + it('returns to the code input when approval fails', async () => { + const dataProvider = { + lookupQuickConnect: vi.fn().mockResolvedValue({ data: finamp }), + authorizeQuickConnect: vi.fn().mockRejectedValue({ status: 409 }), + } + const onClose = renderDialog(dataProvider) + + enterCode('123456') + fireEvent.click(await screen.findByText('menu.quickConnect.approve')) + + await screen.findByRole('textbox') + expect(onClose).not.toHaveBeenCalled() + }) + + it('disables Continue until a code is typed', () => { + renderDialog({ lookupQuickConnect: vi.fn() }) + expect( + screen.getByText('menu.quickConnect.continue').closest('button'), + ).toBeDisabled() + }) +}) diff --git a/ui/src/dialogs/index.js b/ui/src/dialogs/index.js index 86586aef0..3234e4685 100644 --- a/ui/src/dialogs/index.js +++ b/ui/src/dialogs/index.js @@ -2,4 +2,5 @@ export * from './AboutDialog' export * from './SelectPlaylistInput' export * from './ListenBrainzTokenDialog' export * from './SaveQueueDialog' +export * from './QuickConnectDialog' export * from './Dialogs' diff --git a/ui/src/i18n/en.json b/ui/src/i18n/en.json index de96d47c0..f694ea75f 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", @@ -182,6 +183,7 @@ }, "player": { "name": "Player |||| Players", + "menuName": "Players & API keys", "fields": { "name": "Name", "transcodingId": "Transcoding", @@ -190,7 +192,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": { @@ -200,6 +224,9 @@ "targetFormat": "Target Format", "defaultBitRate": "Default Bit Rate", "command": "Command" + }, + "choices": { + "noDefaultBitRate": "None" } }, "playlist": { @@ -304,11 +331,22 @@ "totalDuration": "Duration", "defaultNewUsers": "Default for New Users", "createdAt": "Created", - "updatedAt": "Updated" + "updatedAt": "Updated", + "pidAlbum": "Album grouping", + "pidTrack": "Track identity" }, "sections": { "basic": "Basic Information", - "statistics": "Statistics" + "statistics": "Statistics", + "pid": "Persistent IDs" + }, + "pid": { + "global": "Use global setting (%{value})", + "folder": "Folder (one album per folder)", + "custom": "Custom", + "spec": "PID spec", + "help": "Tags and attributes that identify an item. See the documentation for the syntax:", + "docs": "Persistent IDs" }, "actions": { "scan": "Scan Library", @@ -338,7 +376,9 @@ "messages": { "deleteConfirm": "Are you sure you want to delete this library? This will remove all associated data and user access.", "scanInProgress": "Scan in progress...", - "noLibrariesAssigned": "No libraries assigned to this user" + "noLibrariesAssigned": "No libraries assigned to this user", + "pidChangeTitle": "Change persistent IDs?", + "pidChangeConfirm": "This regroups albums and tracks in this library. A full rescan of this library starts now. Track stars, ratings and play counts are kept. Album stars and ratings move to the new albums where an old album maps to a new one." } }, "plugin": { @@ -588,6 +628,9 @@ "remove_all_missing_content": "Are you sure you want to remove all missing files from the database? This will permanently remove any references to them, including their play counts and ratings.", "notifications_blocked": "You have blocked Notifications for this site in your browser's settings", "notifications_not_available": "This browser does not support desktop notifications or you are not accessing Navidrome over https", + "quickConnectApproved": "%{app} on %{device} is now signed in", + "quickConnectInvalidCode": "Invalid or expired code", + "quickConnectError": "Could not approve the code", "lastfmLinkSuccess": "Last.fm successfully linked and scrobbling enabled", "lastfmLinkFailure": "Last.fm could not be linked", "lastfmUnlinkSuccess": "Last.fm unlinked and scrobbling disabled", @@ -619,6 +662,14 @@ "none": "None" }, "settings": "Settings", + "quickConnect": { + "name": "Quick Connect", + "code": "Code", + "help": "Enter the code shown by a Jellyfin app to sign it in to your account", + "confirm": "Sign in %{app} %{version} on %{device} to your account?", + "continue": "Continue", + "approve": "Approve" + }, "version": "Version", "theme": "Theme", "personal": { diff --git a/ui/src/i18n/useDateLocale.js b/ui/src/i18n/useDateLocale.js new file mode 100644 index 000000000..c9298a473 --- /dev/null +++ b/ui/src/i18n/useDateLocale.js @@ -0,0 +1,14 @@ +import { useLocale } from 'react-admin' + +// Our language codes are mostly region-less ("en"), and Intl reads a bare "en" +// as en-US. Borrow the region from the browser when it speaks the same language. +const resolveDateLocale = (locale, browserLocales = []) => { + if (!locale || locale.includes('-')) return locale + const base = locale.toLowerCase() + return ( + browserLocales.find((l) => l.toLowerCase().split('-')[0] === base) || locale + ) +} + +export const useDateLocale = () => + resolveDateLocale(useLocale(), navigator.languages) diff --git a/ui/src/i18n/useDateLocale.test.js b/ui/src/i18n/useDateLocale.test.js new file mode 100644 index 000000000..c0ea5ab40 --- /dev/null +++ b/ui/src/i18n/useDateLocale.test.js @@ -0,0 +1,45 @@ +import { renderHook } from '@testing-library/react-hooks' +import { describe, it, expect, beforeEach, vi } from 'vitest' +import { useDateLocale } from './useDateLocale' + +vi.mock('react-admin', () => ({ + useLocale: vi.fn(), +})) + +describe('useDateLocale', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + const renderWith = async (locale, browserLocales) => { + const { useLocale } = await import('react-admin') + vi.mocked(useLocale).mockReturnValue(locale) + vi.spyOn(navigator, 'languages', 'get').mockReturnValue(browserLocales) + return renderHook(() => useDateLocale()).result + } + + it('adds the region from the browser when the language has none', async () => { + const result = await renderWith('en', ['en-GB', 'fr-FR']) + expect(result.current).toEqual('en-GB') + }) + + it('falls back to the language when no browser entry matches', async () => { + const result = await renderWith('de', ['en-US', 'fr-FR']) + expect(result.current).toEqual('de') + }) + + it('keeps a language that already carries a region or script', async () => { + expect((await renderWith('pt-br', ['pt-PT'])).current).toEqual('pt-br') + expect((await renderWith('zh-Hans', ['zh-TW'])).current).toEqual('zh-Hans') + }) + + it('matches the browser language case-insensitively', async () => { + const result = await renderWith('pt', ['PT-PT']) + expect(result.current).toEqual('PT-PT') + }) + + it('returns undefined when there is no language', async () => { + const result = await renderWith(undefined, ['en-GB']) + expect(result.current).toBeUndefined() + }) +}) diff --git a/ui/src/layout/AppBar.jsx b/ui/src/layout/AppBar.jsx index 561701dce..460d33bb9 100644 --- a/ui/src/layout/AppBar.jsx +++ b/ui/src/layout/AppBar.jsx @@ -6,12 +6,17 @@ import { usePermissions, getResources, } from 'react-admin' -import { MdInfo, MdPerson, MdSupervisorAccount } from 'react-icons/md' +import { + MdInfo, + MdPerson, + MdPhonelink, + MdSupervisorAccount, +} from 'react-icons/md' import { useSelector } from 'react-redux' import { makeStyles, MenuItem, ListItemIcon, Divider } from '@material-ui/core' import ViewListIcon from '@material-ui/icons/ViewList' import { Dialogs } from '../dialogs/Dialogs' -import { AboutDialog } from '../dialogs' +import { AboutDialog, QuickConnectDialog } from '../dialogs' import PersonalMenu from './PersonalMenu' import ActivityPanel from './ActivityPanel' import NowPlayingPanel from './NowPlayingPanel' @@ -33,33 +38,34 @@ const useStyles = makeStyles( }, ) -const AboutMenuItem = forwardRef(({ onClick, ...rest }, ref) => { - const classes = useStyles(rest) - const translate = useTranslate() - const [open, setOpen] = React.useState(false) +const DialogMenuItem = forwardRef( + ({ onClick, label, icon, dialog, ...rest }, ref) => { + const classes = useStyles(rest) + const [open, setOpen] = React.useState(false) - const handleOpen = () => { - setOpen(true) - } - const handleClose = () => { - onClick && onClick() - setOpen(false) - } - const label = translate('menu.about') - return ( - <> - - - - - {label} - - - - ) -}) + const handleClose = () => { + onClick && onClick() + setOpen(false) + } + return ( + <> + setOpen(true)} + className={classes.root} + > + + {createElement(icon, { title: label, size: 24 })} + + {label} + + {createElement(dialog, { onClose: handleClose, open })} + + ) + }, +) -AboutMenuItem.displayName = 'AboutMenuItem' +DialogMenuItem.displayName = 'DialogMenuItem' const settingsResources = (resource) => resource.name !== 'user' && @@ -96,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 ( { {config.devActivityPanel && permissions === 'admin' && } + {config.enableQuickConnect && ( + + )} {renderUserMenuItemLink()} {resources .filter(settingsResources) .map((r) => renderSettingsMenuItemLink(r))} - + diff --git a/ui/src/layout/AppBar.test.jsx b/ui/src/layout/AppBar.test.jsx index f39dd75cb..cc5ef781d 100644 --- a/ui/src/layout/AppBar.test.jsx +++ b/ui/src/layout/AppBar.test.jsx @@ -9,11 +9,14 @@ import config from '../config' let store +const mocks = vi.hoisted(() => ({ resources: [] })) + vi.mock('react-admin', () => ({ AppBar: ({ userMenu }) =>
{userMenu}
, + MenuItemLink: ({ primaryText }) =>
{primaryText}
, useTranslate: () => (x) => x, usePermissions: () => ({ permissions: 'admin' }), - getResources: () => [], + getResources: () => mocks.resources, })) vi.mock('./NowPlayingPanel', () => ({ @@ -33,12 +36,15 @@ vi.mock('../dialogs/Dialogs', () => ({ })) vi.mock('../dialogs', () => ({ AboutDialog: () =>
, + QuickConnectDialog: () =>
, })) describe('', () => { beforeEach(() => { config.devActivityPanel = true config.enableNowPlaying = true + config.enableQuickConnect = false + mocks.resources = [] store = createStore(combineReducers({ activity: activityReducer }), { activity: { nowPlayingCount: 0 }, }) @@ -62,4 +68,42 @@ describe('', () => { ) expect(screen.queryByTestId('now-playing-panel')).toBeNull() }) + + it('shows the Quick Connect menu item when enabled', () => { + config.enableQuickConnect = true + render( + + + , + ) + expect(screen.queryAllByText('menu.quickConnect.name')).not.toHaveLength(0) + }) + + it('hides the Quick Connect menu item when disabled', () => { + render( + + + , + ) + 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/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) => ( - +// 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 ( + + ) +} + export default Notification diff --git a/ui/src/library/LibraryCreate.jsx b/ui/src/library/LibraryCreate.jsx index 0e69964b6..8166bb2f3 100644 --- a/ui/src/library/LibraryCreate.jsx +++ b/ui/src/library/LibraryCreate.jsx @@ -1,4 +1,5 @@ import React, { useCallback } from 'react' +import PropTypes from 'prop-types' import { Create, SimpleForm, @@ -10,7 +11,34 @@ import { useNotify, useRedirect, } from 'react-admin' +import { Typography } from '@material-ui/core' +import { makeStyles } from '@material-ui/core/styles' import { Title } from '../common' +import { PIDInputs } from './PIDInput' + +const useStyles = makeStyles((theme) => ({ + spaced: { marginTop: theme.spacing(3) }, +})) + +// SimpleForm passes form props (variant, record, ...) to its children, so Typography can't be used directly +const SectionTitle = ({ label, spaced }) => { + const translate = useTranslate() + const classes = useStyles() + return ( + + {translate(label)} + + ) +} + +SectionTitle.propTypes = { + label: PropTypes.string.isRequired, + spaced: PropTypes.bool, +} const LibraryCreate = (props) => { const translate = useTranslate() @@ -73,9 +101,12 @@ const LibraryCreate = (props) => { return ( } {...props}> + + + ) diff --git a/ui/src/library/LibraryEdit.jsx b/ui/src/library/LibraryEdit.jsx index 7e89c892c..c42c7ac4b 100644 --- a/ui/src/library/LibraryEdit.jsx +++ b/ui/src/library/LibraryEdit.jsx @@ -1,12 +1,13 @@ -import React, { useCallback } from 'react' +import React, { useCallback, useState } from 'react' +import PropTypes from 'prop-types' import { Edit, FormWithRedirect, TextInput, BooleanInput, + Confirm, required, SaveButton, - DateField, useTranslate, useMutation, useNotify, @@ -16,8 +17,16 @@ 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' +import config from '../config' +import { PIDInputs } from './PIDInput' +import { pidConfigChanged } from './pidPresets' const useStyles = makeStyles({ toolbar: { @@ -26,6 +35,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 }) @@ -47,8 +58,131 @@ const CustomToolbar = ({ showDelete, ...props }) => ( ) -const LibraryEdit = (props) => { +export const LibraryEditForm = ({ formProps, canEditPath, canDelete }) => { const translate = useTranslate() + const [confirmOpen, setConfirmOpen] = useState(false) + + // Every submit path (Save button and Enter key) goes through here, so a PID change always asks first + const submit = () => { + if ( + pidConfigChanged( + formProps.form.getState().values, + formProps.record, + config, + ) + ) { + setConfirmOpen(true) + return + } + formProps.handleSubmit() + } + + const handleConfirm = () => { + setConfirmOpen(false) + formProps.handleSubmit() + } + + return ( +
{ + event.preventDefault() + submit() + }} + > + + + + {/* Basic Information */} + + {translate('resources.library.sections.basic')} + + + + + + + + + {translate('resources.library.sections.pid')} + + + + + + {/* Statistics - Two Column Layout */} + + {translate('resources.library.sections.statistics')} + + + + + + + + + + + + + + + + + + + + + setConfirmOpen(false)} + /> + + ) +} + +LibraryEditForm.propTypes = { + formProps: PropTypes.object.isRequired, + canEditPath: PropTypes.bool, + canDelete: PropTypes.bool, +} + +const LibraryEdit = (props) => { const [mutate] = useMutation() const notify = useNotify() const redirect = useRedirect() @@ -87,183 +221,11 @@ const LibraryEdit = (props) => { {...props} save={save} render={(formProps) => ( -
- - - - {/* Basic Information */} - - {translate('resources.library.sections.basic')} - - - - - - - - - {/* Statistics - Two Column Layout */} - - {translate('resources.library.sections.statistics')} - - - - - - - - - - - - - - - - - formatBytes(v, 2)} - fullWidth - variant="outlined" - /> - - - - - - - - - - - - - {/* Timestamps Section */} - - - {translate('resources.library.fields.lastScanAt')} - - - - - - - {translate('resources.library.fields.updatedAt')} - - - - - - - {translate('resources.library.fields.createdAt')} - - - - - - - - - + )} /> diff --git a/ui/src/library/LibraryEdit.test.jsx b/ui/src/library/LibraryEdit.test.jsx new file mode 100644 index 000000000..926adc839 --- /dev/null +++ b/ui/src/library/LibraryEdit.test.jsx @@ -0,0 +1,125 @@ +import * as React from 'react' +import { TestContext } from 'ra-test' +import { + FormWithRedirect, + RecordContextProvider, + SaveContextProvider, +} from 'react-admin' +import { + cleanup, + fireEvent, + render, + screen, + waitFor, + within, +} from '@testing-library/react' +import { describe, it, expect, vi, afterEach } from 'vitest' +import { LibraryEditForm } from './LibraryEdit' +import config from '../config' + +const record = { + id: '2', + name: 'Jazz', + path: '/music/jazz', + pidAlbum: '', + pidTrack: '', +} + +// Edit provides a save context in the app. SaveButton only reads these setters from it +const saveContext = { + save: vi.fn(), + setOnSuccess: vi.fn(), + setOnFailure: vi.fn(), + setTransform: vi.fn(), +} + +const renderForm = (save) => + render( + + + + ( + + )} + /> + + + , + ) + +const chooseAlbumGrouping = (optionText) => { + fireEvent.mouseDown( + screen.getByLabelText('resources.library.fields.pidAlbum'), + ) + fireEvent.click(within(screen.getByRole('listbox')).getByText(optionText)) +} + +const dialogTitle = 'resources.library.messages.pidChangeTitle' + +describe('LibraryEditForm', () => { + afterEach(cleanup) + + it('saves directly when the PID config did not change', async () => { + const save = vi.fn() + renderForm(save) + fireEvent.change(screen.getByLabelText(/resources.library.fields.name/), { + target: { value: 'Jazz Renamed' }, + }) + fireEvent.click(screen.getByText('ra.action.save')) + await waitFor(() => expect(save).toHaveBeenCalled()) + expect(screen.queryByText(dialogTitle)).not.toBeInTheDocument() + }) + + it('asks before saving a PID change, and Cancel keeps the edits', async () => { + const save = vi.fn() + renderForm(save) + chooseAlbumGrouping('resources.library.pid.folder') + fireEvent.click(screen.getByText('ra.action.save')) + + expect(await screen.findByText(dialogTitle)).toBeInTheDocument() + expect(save).not.toHaveBeenCalled() + + fireEvent.click(screen.getByText('ra.action.cancel')) + await waitFor(() => + expect(screen.queryByText(dialogTitle)).not.toBeInTheDocument(), + ) + expect(save).not.toHaveBeenCalled() + expect(screen.getByText('resources.library.pid.folder')).toBeInTheDocument() + }) + + it('saves the PID change after Confirm', async () => { + const save = vi.fn() + renderForm(save) + chooseAlbumGrouping('resources.library.pid.folder') + fireEvent.click(screen.getByText('ra.action.save')) + fireEvent.click(await screen.findByText('ra.action.confirm')) + + await waitFor(() => expect(save).toHaveBeenCalled()) + expect(save.mock.calls[0][0]).toMatchObject({ pidAlbum: 'folder' }) + }) + + it('pre-fills a Custom spec with the global spec', () => { + renderForm(vi.fn()) + chooseAlbumGrouping('resources.library.pid.custom') + expect(screen.getByLabelText(/resources.library.pid.spec/)).toHaveValue( + config.pidAlbum, + ) + }) + + it('asks before saving when the form is submitted with Enter', async () => { + const save = vi.fn() + const { container } = renderForm(save) + chooseAlbumGrouping('resources.library.pid.folder') + fireEvent.submit(container.querySelector('form')) + + expect(await screen.findByText(dialogTitle)).toBeInTheDocument() + expect(save).not.toHaveBeenCalled() + }) +}) diff --git a/ui/src/library/PIDInput.jsx b/ui/src/library/PIDInput.jsx new file mode 100644 index 000000000..6481dc392 --- /dev/null +++ b/ui/src/library/PIDInput.jsx @@ -0,0 +1,114 @@ +import React, { useState } from 'react' +import PropTypes from 'prop-types' +import { TextInput, required, useTranslate } from 'react-admin' +import { useField } from 'react-final-form' +import { FormHelperText, Link, MenuItem, TextField } from '@material-ui/core' +import { makeStyles } from '@material-ui/core/styles' +import { + PID_CUSTOM, + PID_FOLDER, + PID_GLOBAL, + pidModeFromValue, + pidValueForMode, +} from './pidPresets' +import config from '../config' +import { docsUrl } from '../utils' + +const PID_DOCS_URL = docsUrl('/docs/usage/pids/') + +const useStyles = makeStyles((theme) => ({ + help: { marginBottom: theme.spacing(1) }, +})) + +// PIDInput edits a library PID override: use the global setting, a preset, or a custom spec +export const PIDInput = ({ source, label, globalValue, allowFolder }) => { + const translate = useTranslate() + const classes = useStyles() + const { input } = useField(source) + // Local state, so choosing Custom shows the text box before anything is typed + const [mode, setMode] = useState(() => + pidModeFromValue(input.value, allowFolder), + ) + + const choices = [ + { + id: PID_GLOBAL, + name: translate('resources.library.pid.global', { value: globalValue }), + }, + ...(allowFolder + ? [{ id: PID_FOLDER, name: translate('resources.library.pid.folder') }] + : []), + { id: PID_CUSTOM, name: translate('resources.library.pid.custom') }, + ] + + const handleModeChange = (event) => { + const newMode = event.target.value + setMode(newMode) + input.onChange(pidValueForMode(newMode, globalValue)) + } + + return ( + <> + + {choices.map((choice) => ( + + {choice.name} + + ))} + + {mode === PID_CUSTOM && ( + <> + + + {translate('resources.library.pid.help')}{' '} + + {translate('resources.library.pid.docs')} + + + + )} + + ) +} + +PIDInput.propTypes = { + source: PropTypes.string.isRequired, + label: PropTypes.string.isRequired, + globalValue: PropTypes.string, + allowFolder: PropTypes.bool, +} + +export const PIDInputs = () => { + const translate = useTranslate() + return ( + <> + + + + ) +} diff --git a/ui/src/library/pidPresets.js b/ui/src/library/pidPresets.js new file mode 100644 index 000000000..0483fc691 --- /dev/null +++ b/ui/src/library/pidPresets.js @@ -0,0 +1,33 @@ +export const PID_GLOBAL = 'global' +export const PID_FOLDER = 'folder' +export const PID_CUSTOM = 'custom' + +export const pidModeFromValue = (value, allowFolder) => { + const v = (value || '').trim() + if (v === '') return PID_GLOBAL + if (allowFolder && v === PID_FOLDER) return PID_FOLDER + return PID_CUSTOM +} + +// Returns the value to store for a dropdown choice. Custom starts from the global spec +export const pidValueForMode = (mode, globalValue) => { + switch (mode) { + case PID_GLOBAL: + return '' + case PID_FOLDER: + return PID_FOLDER + default: + return globalValue || '' + } +} + +// Reports whether the form values change the effective PID spec of the saved record. Like the +// server, it trims, treats empty as the global value and compares case-insensitively +export const pidConfigChanged = (values, record, globals) => { + const effective = (value, field) => + ((value || '').trim() || globals[field] || '').toLowerCase() + return ['pidAlbum', 'pidTrack'].some( + (field) => + effective(values[field], field) !== effective(record[field], field), + ) +} diff --git a/ui/src/library/pidPresets.test.js b/ui/src/library/pidPresets.test.js new file mode 100644 index 000000000..dab1c82e2 --- /dev/null +++ b/ui/src/library/pidPresets.test.js @@ -0,0 +1,65 @@ +import { describe, it, expect } from 'vitest' +import { + PID_CUSTOM, + PID_FOLDER, + PID_GLOBAL, + pidConfigChanged, + pidModeFromValue, + pidValueForMode, +} from './pidPresets' + +describe('pidModeFromValue', () => { + it('maps an empty value to the global setting', () => { + expect(pidModeFromValue('', true)).toBe(PID_GLOBAL) + expect(pidModeFromValue(undefined, true)).toBe(PID_GLOBAL) + }) + it('maps folder to the Folder preset when allowed', () => { + expect(pidModeFromValue('folder', true)).toBe(PID_FOLDER) + }) + it('maps folder to Custom when the Folder preset is not offered', () => { + expect(pidModeFromValue('folder', false)).toBe(PID_CUSTOM) + }) + it('maps any other value to Custom', () => { + expect(pidModeFromValue('album|title', true)).toBe(PID_CUSTOM) + }) +}) + +describe('pidValueForMode', () => { + it('stores an empty value for the global setting', () => { + expect(pidValueForMode(PID_GLOBAL, 'album')).toBe('') + }) + it('stores folder for the Folder preset', () => { + expect(pidValueForMode(PID_FOLDER, '')).toBe('folder') + }) + it('starts Custom from the global spec', () => { + expect(pidValueForMode(PID_CUSTOM, 'album|title')).toBe('album|title') + expect(pidValueForMode(PID_CUSTOM, undefined)).toBe('') + }) +}) + +describe('pidConfigChanged', () => { + const record = { pidAlbum: 'folder', pidTrack: '' } + const globals = { + pidAlbum: 'musicbrainz_albumid|albumartistid,album', + pidTrack: 'musicbrainz_trackid|albumid,discnumber,tracknumber,title', + } + it.each([ + ['nothing changed', { pidAlbum: 'folder', pidTrack: '' }, false], + ['a missing value equals an empty one', { pidAlbum: 'folder' }, false], + [ + 'Custom set to the global value', + { pidAlbum: 'folder', pidTrack: globals.pidTrack }, + false, + ], + ['a case-only change', { pidAlbum: 'FOLDER', pidTrack: '' }, false], + [ + 'a whitespace-only change', + { pidAlbum: ' folder ', pidTrack: ' ' }, + false, + ], + ['the album PID changed', { pidAlbum: '', pidTrack: '' }, true], + ['the track PID changed', { pidAlbum: 'folder', pidTrack: 'title' }, true], + ])('%s', (_, values, expected) => { + expect(pidConfigChanged(values, record, globals)).toBe(expected) + }) +}) 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..d785eb04e 100644 --- a/ui/src/player/PlayerEdit.jsx +++ b/ui/src/player/PlayerEdit.jsx @@ -1,17 +1,16 @@ import { - TextInput, - BooleanInput, - TextField, Edit, - required, SimpleForm, - SelectInput, - ReferenceInput, useTranslate, + DeleteButton, + DeleteWithConfirmButton, + SaveButton, + Toolbar, } from 'react-admin' -import { Title } from '../common' -import config from '../config' -import { BITRATE_CHOICES } from '../consts' +import { makeStyles } from '@material-ui/core/styles' +import { ReadOnlyTextField, Title } from '../common' +import ApiKeyInput from './ApiKeyInput' +import { playerInputs } from './playerInputs' const PlayerTitle = ({ record }) => { const translate = useTranslate() @@ -19,24 +18,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 /> - )} - <TextField source="client" /> - <TextField source="userName" /> + <Edit title={<PlayerTitle />} mutationMode="pessimistic" {...props}> + <SimpleForm variant={'outlined'} toolbar={<PlayerEditToolbar />}> + {playerInputs()} + <ReadOnlyTextField source="client" /> + <ReadOnlyTextField 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) 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/playlist/PlaylistList.jsx b/ui/src/playlist/PlaylistList.jsx index d2b17b108..e1695e980 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,48 +86,25 @@ const TogglePublicInput = ({ resource, source }) => { ) const handleClick = (e) => { - togglePublic() + toggle() e.stopPropagation() } + if (!record) return null + return ( <Switch checked={record[source]} + color="primary" onClick={handleClick} disabled={!isWritable(record.ownerId)} /> ) } -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 +146,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..c05833166 100644 --- a/ui/src/playlist/PlaylistList.test.jsx +++ b/ui/src/playlist/PlaylistList.test.jsx @@ -1,7 +1,9 @@ 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 { RecordContextProvider } from 'react-admin' +import { PlaylistLove, ToggleField, ToggleAutoImport } from './PlaylistList' vi.mock('../config', () => ({ default: { enableFavourites: true }, @@ -13,6 +15,7 @@ vi.mock('../common', () => ({ {record?.starred ? 'starred' : 'not-starred'} </button> ), + isWritable: (ownerId) => ownerId === 'me', })) describe('<PlaylistLove />', () => { @@ -32,3 +35,50 @@ 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('') + }) +}) + +// Secondary is a surface color in many themes, so these toggles must use primary +describe('<ToggleField />', () => { + const renderToggle = (record) => + render( + <TestContext> + <RecordContextProvider value={record}> + <ToggleField resource="playlist" source="public" /> + </RecordContextProvider> + </TestContext>, + ) + + it.each([ + ['owner', 'me', false], + ['non-owner', 'someone-else', true], + ])('renders a primary-colored switch for the %s', (_, ownerId, disabled) => { + renderToggle({ id: 'pl-1', public: true, ownerId }) + const input = screen.getByRole('checkbox') + const switchBase = input.closest('.MuiSwitch-switchBase') + expect(input.checked).toBe(true) + expect(input.disabled).toBe(disabled) + expect(switchBase.classList).toContain('MuiSwitch-colorPrimary') + expect(switchBase.classList).not.toContain('MuiSwitch-colorSecondary') + }) +}) 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/setupTests.js b/ui/src/setupTests.js index ddb999f3c..7cb46e09c 100644 --- a/ui/src/setupTests.js +++ b/ui/src/setupTests.js @@ -14,6 +14,9 @@ const localStorageMock = (function () { setItem: function (key, value) { store[key] = value.toString() }, + removeItem: function (key) { + delete store[key] + }, clear: function () { store = {} }, 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/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/dracula.js b/ui/src/themes/dracula.js index 2e4ae38e5..45559c3af 100644 --- a/ui/src/themes/dracula.js +++ b/ui/src/themes/dracula.js @@ -185,16 +185,6 @@ export default { color: `${foreground} !important`, }, }, - MuiSwitch: { - colorSecondary: { - '&$checked': { - color: green, - }, - '&$checked + $track': { - backgroundColor: green, - }, - }, - }, NDAlbumGridView: { albumName: { marginTop: '0.5rem', diff --git a/ui/src/themes/gruvboxDark.js b/ui/src/themes/gruvboxDark.js index 0f4cbd7c4..3e2955dcd 100644 --- a/ui/src/themes/gruvboxDark.js +++ b/ui/src/themes/gruvboxDark.js @@ -121,16 +121,6 @@ export default { boxShadow: '3px 3px 5px #3c3836', }, }, - MuiSwitch: { - colorSecondary: { - '&$checked': { - color: '#458588', - }, - '&$checked + $track': { - backgroundColor: '#458588', - }, - }, - }, NDMobileArtistDetails: { bgContainer: { background: 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'], diff --git a/ui/src/themes/tokyoNight.js b/ui/src/themes/tokyoNight.js index 07d372a6b..9f6424b77 100644 --- a/ui/src/themes/tokyoNight.js +++ b/ui/src/themes/tokyoNight.js @@ -184,16 +184,6 @@ export default { color: `${foreground} !important`, }, }, - MuiSwitch: { - colorSecondary: { - '&$checked': { - color: blue, - }, - '&$checked + $track': { - backgroundColor: blue, - }, - }, - }, NDAlbumGridView: { albumName: { marginTop: '0.5rem', diff --git a/ui/src/themes/tokyoNightLight.js b/ui/src/themes/tokyoNightLight.js index f84cd0be9..a61c0fe87 100644 --- a/ui/src/themes/tokyoNightLight.js +++ b/ui/src/themes/tokyoNightLight.js @@ -184,16 +184,6 @@ export default { color: `${foreground} !important`, }, }, - MuiSwitch: { - colorSecondary: { - '&$checked': { - color: blue, - }, - '&$checked + $track': { - backgroundColor: blue, - }, - }, - }, NDAlbumGridView: { albumName: { marginTop: '0.5rem', diff --git a/ui/src/themes/useCurrentTheme.js b/ui/src/themes/useCurrentTheme.js index 2b5d13d13..fbb5e9bc8 100644 --- a/ui/src/themes/useCurrentTheme.js +++ b/ui/src/themes/useCurrentTheme.js @@ -59,7 +59,13 @@ const useCurrentTheme = () => { return useMemo( () => ({ ...theme, - props: { ...theme.props, MuiUseMediaQuery: { noSsr: true } }, + props: { + ...theme.props, + MuiUseMediaQuery: { noSsr: true }, + MuiPopover: { disableScrollLock: true }, + // MUI defaults to secondary, which many themes use as a surface color + MuiSwitch: { color: 'primary' }, + }, }), [theme], ) diff --git a/ui/src/themes/useCurrentTheme.test.jsx b/ui/src/themes/useCurrentTheme.test.jsx index 65c3be8c6..6553d9866 100644 --- a/ui/src/themes/useCurrentTheme.test.jsx +++ b/ui/src/themes/useCurrentTheme.test.jsx @@ -3,6 +3,10 @@ import { Provider } from 'react-redux' import { createStore } from 'redux' import mediaQuery from 'css-mediaquery' import { renderHook } from '@testing-library/react-hooks' +import { render, screen } from '@testing-library/react' +import { createMuiTheme, ThemeProvider } from '@material-ui/core/styles' +import Switch from '@material-ui/core/Switch' +import themes from './index' import useCurrentTheme from './useCurrentTheme' import { themeReducer } from '../reducers/themeReducer' import { AUTO_THEME_ID } from '../consts' @@ -161,4 +165,27 @@ describe('useCurrentTheme', () => { expect(document.body.style.backgroundColor).toBe('rgb(18, 18, 18)') }) }) + describe('switch color', () => { + it.each(Object.keys(themes))( + 'renders switches with the primary color in %s', + (theme) => { + const { result } = renderHook(() => useCurrentTheme(), { + wrapper: ({ children }) => ( + <Provider store={createStore(themeReducer, { theme })}> + {children} + </Provider> + ), + }) + render( + <ThemeProvider theme={createMuiTheme(result.current)}> + <Switch checked onChange={() => {}} /> + </ThemeProvider>, + ) + const switchBase = screen + .getByRole('checkbox') + .closest('.MuiSwitch-switchBase') + expect(switchBase.classList).toContain('MuiSwitch-colorPrimary') + }, + ) + }) }) diff --git a/ui/src/transcoding/TranscodingCreate.jsx b/ui/src/transcoding/TranscodingCreate.jsx index aaf122665..014a94e48 100644 --- a/ui/src/transcoding/TranscodingCreate.jsx +++ b/ui/src/transcoding/TranscodingCreate.jsx @@ -8,7 +8,7 @@ import { useTranslate, } from 'react-admin' import { Title } from '../common' -import { BITRATE_CHOICES } from '../consts' +import { TRANSCODING_BITRATE_CHOICES } from '../consts' const TranscodingTitle = () => { const translate = useTranslate() @@ -28,7 +28,7 @@ const TranscodingCreate = (props) => ( <TextInput source="targetFormat" validate={[required()]} /> <SelectInput source="defaultBitRate" - choices={BITRATE_CHOICES} + choices={TRANSCODING_BITRATE_CHOICES} defaultValue={192} /> <TextInput diff --git a/ui/src/transcoding/TranscodingEdit.jsx b/ui/src/transcoding/TranscodingEdit.jsx index 5aba9a4e3..bf500920f 100644 --- a/ui/src/transcoding/TranscodingEdit.jsx +++ b/ui/src/transcoding/TranscodingEdit.jsx @@ -9,7 +9,7 @@ import { } from 'react-admin' import { Title } from '../common' import { TranscodingNote } from './TranscodingNote' -import { BITRATE_CHOICES } from '../consts' +import { TRANSCODING_BITRATE_CHOICES } from '../consts' const TranscodingTitle = ({ record }) => { const translate = useTranslate() @@ -28,7 +28,10 @@ const TranscodingEdit = (props) => { <SimpleForm variant={'outlined'}> <TextInput source="name" validate={[required()]} /> <TextInput source="targetFormat" validate={[required()]} /> - <SelectInput source="defaultBitRate" choices={BITRATE_CHOICES} /> + <SelectInput + source="defaultBitRate" + choices={TRANSCODING_BITRATE_CHOICES} + /> <TextInput source="command" fullWidth validate={[required()]} /> </SimpleForm> </Edit> diff --git a/ui/src/transcoding/TranscodingList.jsx b/ui/src/transcoding/TranscodingList.jsx index bca8b49df..d1371c04a 100644 --- a/ui/src/transcoding/TranscodingList.jsx +++ b/ui/src/transcoding/TranscodingList.jsx @@ -1,7 +1,8 @@ import React from 'react' -import { Datagrid, TextField } from 'react-admin' +import { Datagrid, SelectField, TextField } from 'react-admin' import { useMediaQuery } from '@material-ui/core' import { SimpleList, List } from '../common' +import { TRANSCODING_BITRATE_CHOICES } from '../consts' import config from '../config' const TranscodingList = (props) => { @@ -16,13 +17,22 @@ const TranscodingList = (props) => { <SimpleList primaryText={(r) => r.name} secondaryText={(r) => `format: ${r.targetFormat}`} - tertiaryText={(r) => r.defaultBitRate} + tertiaryText={(r) => ( + <SelectField + record={r} + source="defaultBitRate" + choices={TRANSCODING_BITRATE_CHOICES} + /> + )} /> ) : ( <Datagrid rowClick={config.enableTranscodingConfig ? 'edit' : 'show'}> <TextField source="name" /> <TextField source="targetFormat" /> - <TextField source="defaultBitRate" /> + <SelectField + source="defaultBitRate" + choices={TRANSCODING_BITRATE_CHOICES} + /> <TextField source="command" /> </Datagrid> )} diff --git a/ui/src/transcoding/TranscodingShow.jsx b/ui/src/transcoding/TranscodingShow.jsx index e132afee5..b7ec2f595 100644 --- a/ui/src/transcoding/TranscodingShow.jsx +++ b/ui/src/transcoding/TranscodingShow.jsx @@ -1,7 +1,8 @@ import React from 'react' -import { Show, SimpleShowLayout, TextField } from 'react-admin' +import { SelectField, Show, SimpleShowLayout, TextField } from 'react-admin' import { Title } from '../common' import { TranscodingNote } from './TranscodingNote' +import { TRANSCODING_BITRATE_CHOICES } from '../consts' const TranscodingTitle = ({ record }) => { return <Title subTitle={`Transcoding ${record ? record.name : ''}`} /> @@ -16,7 +17,10 @@ const TranscodingShow = (props) => { <SimpleShowLayout> <TextField source="name" /> <TextField source="targetFormat" /> - <TextField source="defaultBitRate" /> + <SelectField + source="defaultBitRate" + choices={TRANSCODING_BITRATE_CHOICES} + /> <TextField source="command" /> </SimpleShowLayout> </Show> 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 diff --git a/ui/src/user/UserList.jsx b/ui/src/user/UserList.jsx index 4faf31785..7922ca08d 100644 --- a/ui/src/user/UserList.jsx +++ b/ui/src/user/UserList.jsx @@ -9,6 +9,8 @@ import { } from 'react-admin' import { useMediaQuery } from '@material-ui/core' import { List, DateField } from '../common' +import { useDateLocale } from '../i18n/useDateLocale' +import { formatDateTime } from '../utils/formatters' const UserFilter = (props) => ( <Filter {...props} variant={'outlined'}> @@ -18,6 +20,7 @@ const UserFilter = (props) => ( const UserList = (props) => { const isXsmall = useMediaQuery((theme) => theme.breakpoints.down('xs')) + const locale = useDateLocale() return ( <List @@ -31,7 +34,7 @@ const UserList = (props) => { <SimpleList primaryText={(record) => record.userName} secondaryText={(record) => - record.lastLoginAt && new Date(record.lastLoginAt).toLocaleString() + record.lastLoginAt && formatDateTime(record.lastLoginAt, locale) } tertiaryText={(record) => (record.isAdmin ? '[admin]️' : '')} /> diff --git a/ui/src/utils/formatters.js b/ui/src/utils/formatters.js index cfcb84b05..f7c59c195 100644 --- a/ui/src/utils/formatters.js +++ b/ui/src/utils/formatters.js @@ -95,6 +95,9 @@ export const formatFullDate = (date, locale) => { return new Date(date).toLocaleDateString(locale, options) } +export const formatDateTime = (value, locale) => + new Date(value).toLocaleString(locale) + export const formatNumber = (value, locale) => { if (value === null || value === undefined) return '0' return value.toLocaleString(locale) diff --git a/utils/httpclient/httpclient.go b/utils/httpclient/httpclient.go index 7fb48f36d..b0e7b5681 100644 --- a/utils/httpclient/httpclient.go +++ b/utils/httpclient/httpclient.go @@ -3,10 +3,15 @@ package httpclient import ( + "context" + "net" "net/http" + "net/netip" + "net/url" "time" "github.com/navidrome/navidrome/consts" + "github.com/navidrome/navidrome/utils/netguard" ) type uaTransport struct { @@ -33,3 +38,47 @@ func NewTransport(base http.RoundTripper) http.RoundTripper { func New(timeout time.Duration) *http.Client { return &http.Client{Timeout: timeout, Transport: NewTransport(nil)} } + +// proxyFunc resolves the proxy for a request; tests replace it. +var proxyFunc = http.ProxyFromEnvironment + +type proxyAddrKey struct{} + +// NewExternal is New for URLs from untrusted sources: it refuses to dial private, loopback, +// link-local and unspecified addresses, except those covered by allowed. +func NewExternal(timeout time.Duration, allowed ...netip.Prefix) *http.Client { + t := http.DefaultTransport.(*http.Transport).Clone() + t.Proxy = proxyFunc + direct := net.Dialer{Timeout: 30 * time.Second, KeepAlive: 30 * time.Second} + guarded := direct + guarded.Control = netguard.DialControl(allowed...) + t.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + // Exempt the hop to the proxy, which relays the request and is operator config. A URL + // aimed at the proxy's own address is not proxied, so it stays guarded. + if proxy, ok := ctx.Value(proxyAddrKey{}).(string); ok && proxy == addr { + return direct.DialContext(ctx, network, addr) + } + return guarded.DialContext(ctx, network, addr) + } + return &http.Client{Timeout: timeout, Transport: NewTransport(&proxyTagger{base: t})} +} + +// proxyTagger records the proxy each request resolves to, so the dialer can tell a hop to the +// proxy from a dial to the URL's own host. +type proxyTagger struct{ base http.RoundTripper } + +func (p *proxyTagger) RoundTrip(req *http.Request) (*http.Response, error) { + if u, err := proxyFunc(req); err == nil && u != nil { + req = req.WithContext(context.WithValue(req.Context(), proxyAddrKey{}, proxyAddr(u))) + } + return p.base.RoundTrip(req) +} + +// proxyAddr mirrors how net/http addresses a proxy connection. +func proxyAddr(u *url.URL) string { + port := u.Port() + if port == "" { + port = map[string]string{"http": "80", "https": "443", "socks5": "1080", "socks5h": "1080"}[u.Scheme] + } + return net.JoinHostPort(u.Hostname(), port) +} diff --git a/utils/httpclient/httpclient_proxy_internal_test.go b/utils/httpclient/httpclient_proxy_internal_test.go new file mode 100644 index 000000000..2d542d8fe --- /dev/null +++ b/utils/httpclient/httpclient_proxy_internal_test.go @@ -0,0 +1,68 @@ +package httpclient + +import ( + "net/http" + "net/http/httptest" + "net/url" + "sync/atomic" + "time" + + "github.com/navidrome/navidrome/utils/netguard" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("NewExternal with a proxy", func() { + var proxy *httptest.Server + var proxied atomic.Int32 + var proxyURL *url.URL + + BeforeEach(func() { + proxied.Store(0) + proxy = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + proxied.Add(1) + _, _ = w.Write([]byte("via proxy")) + })) + DeferCleanup(proxy.Close) + proxyURL, _ = url.Parse(proxy.URL) + prev := proxyFunc + DeferCleanup(func() { proxyFunc = prev }) + }) + + It("dials the configured proxy even though it listens on a private address", func() { + proxyFunc = func(*http.Request) (*url.URL, error) { return proxyURL, nil } + + resp, err := NewExternal(time.Second).Get("http://navidrome.example.com/cover.jpg") + Expect(err).ToNot(HaveOccurred()) + defer resp.Body.Close() + Expect(proxied.Load()).To(Equal(int32(1))) + }) + + // net/http never proxies loopback targets, so a URL aimed at a loopback proxy is dialed directly. + It("refuses a direct dial to the proxy's own address when the request is not proxied", func() { + proxyFunc = func(r *http.Request) (*url.URL, error) { + if r.URL.Hostname() == "127.0.0.1" { + return nil, nil + } + return proxyURL, nil + } + + _, err := NewExternal(time.Second).Get(proxy.URL + "/secret") + Expect(err).To(MatchError(netguard.ErrPrivateAddress)) + Expect(proxied.Load()).To(BeZero()) + }) + + It("still refuses a direct dial to a private address", func() { + proxyFunc = func(r *http.Request) (*url.URL, error) { + if r.URL.Scheme == "https" { + return proxyURL, nil + } + return nil, nil + } + target := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + DeferCleanup(target.Close) + + _, err := NewExternal(time.Second).Get(target.URL) + Expect(err).To(MatchError(netguard.ErrPrivateAddress)) + }) +}) diff --git a/utils/httpclient/httpclient_test.go b/utils/httpclient/httpclient_test.go index c86b51165..7b1e58797 100644 --- a/utils/httpclient/httpclient_test.go +++ b/utils/httpclient/httpclient_test.go @@ -3,10 +3,12 @@ package httpclient_test import ( "net/http" "net/http/httptest" + "net/netip" "time" "github.com/navidrome/navidrome/consts" "github.com/navidrome/navidrome/utils/httpclient" + "github.com/navidrome/navidrome/utils/netguard" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -16,6 +18,7 @@ var _ = Describe("httpclient", func() { var receivedUA string BeforeEach(func() { + receivedUA = "" server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { receivedUA = r.Header.Get("User-Agent") })) @@ -48,6 +51,34 @@ var _ = Describe("httpclient", func() { }) }) + Describe("NewExternal", func() { + It("refuses to connect to a loopback server", func() { + c := httpclient.NewExternal(time.Second) + _, err := c.Get(server.URL) + Expect(err).To(MatchError(netguard.ErrPrivateAddress)) + Expect(receivedUA).To(BeEmpty()) + }) + + It("refuses a redirect from an allowed host to a private address", func() { + redirector := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound) + })) + DeferCleanup(redirector.Close) + + c := httpclient.NewExternal(time.Second, netip.MustParsePrefix("127.0.0.1/32")) + _, err := c.Get(redirector.URL) + Expect(err).To(MatchError(netguard.ErrPrivateAddress)) + }) + + It("connects to addresses covered by an allowed prefix and sets the User-Agent", func() { + c := httpclient.NewExternal(time.Second, netip.MustParsePrefix("127.0.0.0/8")) + resp, err := c.Get(server.URL) + Expect(err).ToNot(HaveOccurred()) + resp.Body.Close() + Expect(receivedUA).To(Equal(consts.HTTPUserAgent)) + }) + }) + Describe("NewTransport", func() { It("uses the default transport when base is nil", func() { c := &http.Client{Transport: httpclient.NewTransport(nil)} diff --git a/utils/netguard/netguard.go b/utils/netguard/netguard.go new file mode 100644 index 000000000..43d80ecc9 --- /dev/null +++ b/utils/netguard/netguard.go @@ -0,0 +1,38 @@ +// Package netguard keeps outbound connections driven by untrusted input away from internal addresses. +package netguard + +import ( + "errors" + "fmt" + "net" + "net/netip" + "syscall" +) + +var ErrPrivateAddress = errors.New("dial to private/loopback address blocked") + +// IsPrivateIP reports whether ip is loopback, private, link-local or unspecified. +func IsPrivateIP(ip net.IP) bool { + return ip.IsLoopback() || ip.IsUnspecified() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() +} + +// DialControl returns a net.Dialer Control hook that rejects private addresses not covered by allowed. +// It sees the resolved IP, so DNS names, redirects and rebinding cannot get around it. +func DialControl(allowed ...netip.Prefix) func(network, address string, c syscall.RawConn) error { + return func(_, address string, _ syscall.RawConn) error { + ap, err := netip.ParseAddrPort(address) + if err != nil { + return fmt.Errorf("%w: unparseable address %q: %w", ErrPrivateAddress, address, err) + } + addr := ap.Addr().Unmap() + if !IsPrivateIP(addr.AsSlice()) { + return nil + } + for _, p := range allowed { + if p.Contains(addr) { + return nil + } + } + return fmt.Errorf("%w: %s", ErrPrivateAddress, address) + } +} diff --git a/scanner/metadata_old/metadata_suite_test.go b/utils/netguard/netguard_suite_test.go similarity index 66% rename from scanner/metadata_old/metadata_suite_test.go rename to utils/netguard/netguard_suite_test.go index 03ec3c847..7769e0c82 100644 --- a/scanner/metadata_old/metadata_suite_test.go +++ b/utils/netguard/netguard_suite_test.go @@ -1,4 +1,4 @@ -package metadata_old +package netguard_test import ( "testing" @@ -9,9 +9,9 @@ import ( . "github.com/onsi/gomega" ) -func TestMetadata(t *testing.T) { - tests.Init(t, true) +func TestNetguard(t *testing.T) { + tests.Init(t, false) log.SetLevel(log.LevelFatal) RegisterFailHandler(Fail) - RunSpecs(t, "Metadata Suite") + RunSpecs(t, "Netguard Suite") } diff --git a/utils/netguard/netguard_test.go b/utils/netguard/netguard_test.go new file mode 100644 index 000000000..733204226 --- /dev/null +++ b/utils/netguard/netguard_test.go @@ -0,0 +1,65 @@ +package netguard_test + +import ( + "net" + "net/netip" + + "github.com/navidrome/navidrome/utils/netguard" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("netguard", func() { + DescribeTable("IsPrivateIP", + func(addr string, expected bool) { + Expect(netguard.IsPrivateIP(net.ParseIP(addr))).To(Equal(expected)) + }, + Entry("IPv4 loopback", "127.0.0.1", true), + Entry("IPv4 loopback range", "127.0.0.2", true), + Entry("IPv6 loopback", "::1", true), + Entry("IPv4-mapped IPv6 loopback", "::ffff:127.0.0.1", true), + Entry("10.x", "10.1.2.3", true), + Entry("172.16.x", "172.16.0.1", true), + Entry("192.168.x", "192.168.1.10", true), + Entry("IPv6 unique local", "fd00::1", true), + Entry("link-local (cloud metadata)", "169.254.169.254", true), + Entry("IPv6 link-local", "fe80::1", true), + Entry("link-local multicast", "224.0.0.1", true), + Entry("IPv4 unspecified", "0.0.0.0", true), + Entry("IPv6 unspecified", "::", true), + Entry("public IPv4", "93.184.216.34", false), + Entry("172.32.x is outside the private block", "172.32.0.1", false), + Entry("public IPv6", "2606:4700:4700::1111", false), + ) + + Describe("DialControl", func() { + It("rejects private and loopback addresses", func() { + control := netguard.DialControl() + for _, addr := range []string{"127.0.0.1:80", "[::1]:443", "169.254.169.254:80", "10.0.0.1:8080", "0.0.0.0:80"} { + Expect(control("tcp", addr, nil)).To(MatchError(netguard.ErrPrivateAddress), addr) + } + }) + + It("allows public addresses", func() { + control := netguard.DialControl() + Expect(control("tcp", "93.184.216.34:443", nil)).To(Succeed()) + Expect(control("tcp6", "[2606:4700:4700::1111]:443", nil)).To(Succeed()) + }) + + It("allows private addresses covered by an allowed prefix, and only those", func() { + control := netguard.DialControl(netip.MustParsePrefix("127.0.0.1/32")) + Expect(control("tcp", "127.0.0.1:8080", nil)).To(Succeed()) + Expect(control("tcp", "127.0.0.2:8080", nil)).To(MatchError(netguard.ErrPrivateAddress)) + }) + + It("matches IPv4-mapped IPv6 addresses against IPv4 prefixes", func() { + control := netguard.DialControl(netip.MustParsePrefix("127.0.0.0/8")) + Expect(control("tcp", "[::ffff:127.0.0.1]:80", nil)).To(Succeed()) + }) + + It("fails closed when the address is not an IP", func() { + Expect(netguard.DialControl()("tcp", "localhost:80", nil)).To(HaveOccurred()) + Expect(netguard.DialControl()("tcp", "garbage", nil)).To(HaveOccurred()) + }) + }) +}) 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), + ) +}) 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) } 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")) + }) + }) +})