From 9be5bbae3cfe41d8484da8b8cf657b782c6a9b16 Mon Sep 17 00:00:00 2001 From: Bryan Van Deusen Date: Thu, 8 Oct 2026 09:24:20 -0400 Subject: [PATCH] fix(discover): cap the taste-matched arm per album and artist before its LIMIT (#5356) Summed tag weight rewards a track for carrying many of the user's tags, so on the deploy two artists whose every track carries the whole lo-fi profile took all 120 rows of the taste-unheard query. capByAlbumAndArtist ran after the LIMIT and left 6, and the arm with the lowest skip rate (12% against ~23%) handed its slots to dormant and random. The query now ranks within album, then within artist over what the album cap kept, before the LIMIT: the same walk the Go cap makes, so the bucket fills from as many artists as match. The caps come from the Go constants. Co-Authored-By: Claude Opus 5.5 --- internal/db/dbq/discover.sql.go | 86 +++++++++++----- internal/db/queries/discover.sql | 73 +++++++++----- internal/playlists/discover.go | 5 +- internal/playlists/discover_taste_db_test.go | 100 +++++++++++++++++++ 4 files changed, 212 insertions(+), 52 deletions(-) create mode 100644 internal/playlists/discover_taste_db_test.go diff --git a/internal/db/dbq/discover.sql.go b/internal/db/dbq/discover.sql.go index 7abd362a..3689aeba 100644 --- a/internal/db/dbq/discover.sql.go +++ b/internal/db/dbq/discover.sql.go @@ -216,35 +216,55 @@ func (q *Queries) ListRandomUnheardTracksForDiscover(ctx context.Context, arg Li } const listTasteUnheardTracksForDiscover = `-- name: ListTasteUnheardTracksForDiscover :many -SELECT t.id, t.album_id, t.artist_id - FROM tracks t - JOIN LATERAL regexp_split_to_table(coalesce(t.genre, ''), '[;,]') AS g_split(g) ON true - JOIN taste_profile_tags nt ON nt.user_id = $1 AND trim(g_split.g) = nt.tag - WHERE t.missing_since IS NULL -- #2523: never offer a file that is gone - AND nt.weight > 0 - AND trim(g_split.g) <> '' - AND NOT EXISTS ( - SELECT 1 FROM play_events pe - WHERE pe.user_id = $1 - AND pe.track_id = t.id - AND pe.was_skipped = false - ) - AND NOT EXISTS ( - SELECT 1 FROM general_likes gl - WHERE gl.user_id = $1 AND gl.track_id = t.id - ) - AND NOT EXISTS ( - SELECT 1 FROM lidarr_quarantine q - WHERE q.user_id = $1 AND q.track_id = t.id - ) - GROUP BY t.id, t.album_id, t.artist_id - ORDER BY SUM(nt.weight) DESC, md5(t.id::text || $2::text) +WITH scored AS ( + SELECT t.id, t.album_id, t.artist_id, + SUM(nt.weight) AS weight, + md5(t.id::text || $2::text) AS tiebreak + FROM tracks t + JOIN LATERAL regexp_split_to_table(coalesce(t.genre, ''), '[;,]') AS g_split(g) ON true + JOIN taste_profile_tags nt ON nt.user_id = $3 AND trim(g_split.g) = nt.tag + WHERE t.missing_since IS NULL -- #2523: never offer a file that is gone + AND nt.weight > 0 + AND trim(g_split.g) <> '' + AND NOT EXISTS ( + SELECT 1 FROM play_events pe + WHERE pe.user_id = $3 + AND pe.track_id = t.id + AND pe.was_skipped = false + ) + AND NOT EXISTS ( + SELECT 1 FROM general_likes gl + WHERE gl.user_id = $3 AND gl.track_id = t.id + ) + AND NOT EXISTS ( + SELECT 1 FROM lidarr_quarantine q + WHERE q.user_id = $3 AND q.track_id = t.id + ) + GROUP BY t.id, t.album_id, t.artist_id +), +album_capped AS ( + SELECT s.id, s.album_id, s.artist_id, s.weight, s.tiebreak, + row_number() OVER (PARTITION BY s.album_id ORDER BY s.weight DESC, s.tiebreak) AS album_rank + FROM scored s +), +artist_capped AS ( + SELECT a.id, a.album_id, a.artist_id, a.weight, a.tiebreak, a.album_rank, + row_number() OVER (PARTITION BY a.artist_id ORDER BY a.weight DESC, a.tiebreak) AS artist_rank + FROM album_capped a + WHERE a.album_rank <= $4::int +) +SELECT c.id, c.album_id, c.artist_id + FROM artist_capped c + WHERE c.artist_rank <= $1::int + ORDER BY c.weight DESC, c.tiebreak LIMIT 120 ` type ListTasteUnheardTracksForDiscoverParams struct { - UserID pgtype.UUID - Column2 string + MaxPerArtist int32 + DateSeed string + UserID pgtype.UUID + MaxPerAlbum int32 } type ListTasteUnheardTracksForDiscoverRow struct { @@ -261,9 +281,21 @@ type ListTasteUnheardTracksForDiscoverRow struct { // [;,]). Same exclusion filters as the other buckets. Returns nothing // when the user has no taste tags yet (cold start), so the caller // redistributes its slots to the other buckets. Stamped 'taste_unheard'. -// $1 = user_id, $2 = date string for md5 tiebreak ordering. +// +// The per-album and per-artist caps apply BEFORE the LIMIT (#5356). Summed +// weight rewards a track for carrying many of the user's tags, so a few +// artists whose every track is tagged with the whole profile take every row +// of a plain LIMIT; the caller's caps then left 6 of 120 on the deploy, and +// the best-performing arm handed its slots to the others. Ranking within +// album, then within artist over what the album cap kept, is the same walk +// capByAlbumAndArtist makes, so the caller's caps keep everything here. func (q *Queries) ListTasteUnheardTracksForDiscover(ctx context.Context, arg ListTasteUnheardTracksForDiscoverParams) ([]ListTasteUnheardTracksForDiscoverRow, error) { - rows, err := q.db.Query(ctx, listTasteUnheardTracksForDiscover, arg.UserID, arg.Column2) + rows, err := q.db.Query(ctx, listTasteUnheardTracksForDiscover, + arg.MaxPerArtist, + arg.DateSeed, + arg.UserID, + arg.MaxPerAlbum, + ) if err != nil { return nil, err } diff --git a/internal/db/queries/discover.sql b/internal/db/queries/discover.sql index 86385609..2c297610 100644 --- a/internal/db/queries/discover.sql +++ b/internal/db/queries/discover.sql @@ -115,28 +115,53 @@ SELECT t.id, t.album_id, t.artist_id -- [;,]). Same exclusion filters as the other buckets. Returns nothing -- when the user has no taste tags yet (cold start), so the caller -- redistributes its slots to the other buckets. Stamped 'taste_unheard'. --- $1 = user_id, $2 = date string for md5 tiebreak ordering. -SELECT t.id, t.album_id, t.artist_id - FROM tracks t - JOIN LATERAL regexp_split_to_table(coalesce(t.genre, ''), '[;,]') AS g_split(g) ON true - JOIN taste_profile_tags nt ON nt.user_id = $1 AND trim(g_split.g) = nt.tag - WHERE t.missing_since IS NULL -- #2523: never offer a file that is gone - AND nt.weight > 0 - AND trim(g_split.g) <> '' - AND NOT EXISTS ( - SELECT 1 FROM play_events pe - WHERE pe.user_id = $1 - AND pe.track_id = t.id - AND pe.was_skipped = false - ) - AND NOT EXISTS ( - SELECT 1 FROM general_likes gl - WHERE gl.user_id = $1 AND gl.track_id = t.id - ) - AND NOT EXISTS ( - SELECT 1 FROM lidarr_quarantine q - WHERE q.user_id = $1 AND q.track_id = t.id - ) - GROUP BY t.id, t.album_id, t.artist_id - ORDER BY SUM(nt.weight) DESC, md5(t.id::text || $2::text) +-- +-- The per-album and per-artist caps apply BEFORE the LIMIT (#5356). Summed +-- weight rewards a track for carrying many of the user's tags, so a few +-- artists whose every track is tagged with the whole profile take every row +-- of a plain LIMIT; the caller's caps then left 6 of 120 on the deploy, and +-- the best-performing arm handed its slots to the others. Ranking within +-- album, then within artist over what the album cap kept, is the same walk +-- capByAlbumAndArtist makes, so the caller's caps keep everything here. +WITH scored AS ( + SELECT t.id, t.album_id, t.artist_id, + SUM(nt.weight) AS weight, + md5(t.id::text || sqlc.arg(date_seed)::text) AS tiebreak + FROM tracks t + JOIN LATERAL regexp_split_to_table(coalesce(t.genre, ''), '[;,]') AS g_split(g) ON true + JOIN taste_profile_tags nt ON nt.user_id = sqlc.arg(user_id) AND trim(g_split.g) = nt.tag + WHERE t.missing_since IS NULL -- #2523: never offer a file that is gone + AND nt.weight > 0 + AND trim(g_split.g) <> '' + AND NOT EXISTS ( + SELECT 1 FROM play_events pe + WHERE pe.user_id = sqlc.arg(user_id) + AND pe.track_id = t.id + AND pe.was_skipped = false + ) + AND NOT EXISTS ( + SELECT 1 FROM general_likes gl + WHERE gl.user_id = sqlc.arg(user_id) AND gl.track_id = t.id + ) + AND NOT EXISTS ( + SELECT 1 FROM lidarr_quarantine q + WHERE q.user_id = sqlc.arg(user_id) AND q.track_id = t.id + ) + GROUP BY t.id, t.album_id, t.artist_id +), +album_capped AS ( + SELECT s.*, + row_number() OVER (PARTITION BY s.album_id ORDER BY s.weight DESC, s.tiebreak) AS album_rank + FROM scored s +), +artist_capped AS ( + SELECT a.*, + row_number() OVER (PARTITION BY a.artist_id ORDER BY a.weight DESC, a.tiebreak) AS artist_rank + FROM album_capped a + WHERE a.album_rank <= sqlc.arg(max_per_album)::int +) +SELECT c.id, c.album_id, c.artist_id + FROM artist_capped c + WHERE c.artist_rank <= sqlc.arg(max_per_artist)::int + ORDER BY c.weight DESC, c.tiebreak LIMIT 120; diff --git a/internal/playlists/discover.go b/internal/playlists/discover.go index f70c658d..e30d3c31 100644 --- a/internal/playlists/discover.go +++ b/internal/playlists/discover.go @@ -95,8 +95,11 @@ type discoverPools struct { } func loadDiscoverPools(ctx context.Context, q *dbq.Queries, logger *slog.Logger, userID pgtype.UUID, dateStr string) discoverPools { + // The caps go into the query too (#5356): applied only here, after its + // LIMIT, they left a heavily tagged artist or two owning the whole bucket. tasteRows, err := q.ListTasteUnheardTracksForDiscover(ctx, dbq.ListTasteUnheardTracksForDiscoverParams{ - UserID: userID, Column2: dateStr, + UserID: userID, DateSeed: dateStr, + MaxPerAlbum: discoverMaxTracksPerAlbum, MaxPerArtist: discoverMaxTracksPerArtist, }) if err != nil { logger.Warn("discover: taste-unheard bucket failed; continuing with empty pool", diff --git a/internal/playlists/discover_taste_db_test.go b/internal/playlists/discover_taste_db_test.go new file mode 100644 index 00000000..81772972 --- /dev/null +++ b/internal/playlists/discover_taste_db_test.go @@ -0,0 +1,100 @@ +package playlists_test + +import ( + "context" + "fmt" + "testing" + + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + "git.fabledsword.com/bvandeusen/minstrel/internal/db/dbq" +) + +func setGenre(t *testing.T, pool *pgxpool.Pool, trackID pgtype.UUID, genre string) { + t.Helper() + if _, err := pool.Exec(context.Background(), + `UPDATE tracks SET genre = $2 WHERE id = $1`, trackID, genre); err != nil { + t.Fatalf("set genre: %v", err) + } +} + +func seedTasteTag(t *testing.T, pool *pgxpool.Pool, userID pgtype.UUID, tag string, weight float64) { + t.Helper() + if _, err := pool.Exec(context.Background(), + `INSERT INTO taste_profile_tags (user_id, tag, weight) VALUES ($1, $2, $3) + ON CONFLICT (user_id, tag) DO UPDATE SET weight = EXCLUDED.weight`, + userID, tag, weight); err != nil { + t.Fatalf("seed taste tag: %v", err) + } +} + +// TestListTasteUnheardTracksForDiscover_CapsBeforeLimit is the deploy's shape +// (#5356): one artist whose every track carries the whole taste profile +// outscores everything on summed weight. Capped only after the LIMIT, that +// artist took all 120 rows and the bucket shrank to its 3. Capped before, the +// artist keeps 3 and the single-tag artists below it fill the rest. +func TestListTasteUnheardTracksForDiscover_CapsBeforeLimit(t *testing.T) { + pool := newPool(t) + ctx := context.Background() + u := seedUser(t, pool, "tastecap") + + profile := []string{"Lo-Fi", "Downtempo", "Hip Hop", "Instrumental", "Chillwave"} + for _, tag := range profile { + seedTasteTag(t, pool, u.ID, tag, 10) + } + allTags := "Lo-Fi; Downtempo; Hip Hop; Instrumental; Chillwave" + + // 13 albums x 10 tracks = 130, more than the query's LIMIT of 120. + heavy := seedTrack(t, pool, "Heavy 0", "Heavy Artist") + setGenre(t, pool, heavy.ID, allTags) + album := heavy.AlbumID + for i := 1; i < 130; i++ { + if i%10 == 0 { + album = seedAlbumForArtist(t, pool, fmt.Sprintf("Heavy Album %d", i/10), heavy.ArtistID) + } + tr := seedTrackForArtist(t, pool, fmt.Sprintf("Heavy %d", i), album, heavy.ArtistID) + setGenre(t, pool, tr.ID, allTags) + } + + // Twelve artists with one track each, matching a single profile tag. + const light = 12 + for i := 0; i < light; i++ { + tr := seedTrack(t, pool, fmt.Sprintf("Light %d", i), fmt.Sprintf("Light Artist %d", i)) + setGenre(t, pool, tr.ID, profile[i%len(profile)]) + } + + rows, err := dbq.New(pool).ListTasteUnheardTracksForDiscover(ctx, dbq.ListTasteUnheardTracksForDiscoverParams{ + UserID: u.ID, DateSeed: "2026-10-08", MaxPerAlbum: 2, MaxPerArtist: 3, + }) + if err != nil { + t.Fatalf("list taste unheard: %v", err) + } + + perArtist := map[pgtype.UUID]int{} + perAlbum := map[pgtype.UUID]int{} + for _, r := range rows { + perArtist[r.ArtistID]++ + perAlbum[r.AlbumID]++ + } + if got := perArtist[heavy.ArtistID]; got != 3 { + t.Errorf("heavy artist rows = %d, want exactly the artist cap of 3", got) + } + for id, n := range perAlbum { + if n > 2 { + t.Errorf("album %v has %d rows, over the album cap of 2", id, n) + } + } + if got, want := len(perArtist), 1+light; got != want { + t.Errorf("distinct artists = %d, want %d (the heavy artist plus every single-tag artist)", got, want) + } + if got, want := len(rows), 3+light; got != want { + t.Errorf("rows = %d, want %d", got, want) + } + // Still ranked by weight: the heavy artist's three lead. + for i := 0; i < 3 && i < len(rows); i++ { + if rows[i].ArtistID != heavy.ArtistID { + t.Errorf("row %d is not the heavy artist's; higher summed weight should rank first", i) + } + } +}