Files
minstrel/internal/playlists/discover_taste_db_test.go
T
bvandeusenandClaude Opus 5.5 9be5bbae3c
release / web (push) Successful in 1m37s
release / govulncheck (push) Successful in 53s
release / go (push) Successful in 2m7s
release / integration (push) Successful in 5m2s
release / android (push) Successful in 6m34s
release / Build signed APK (releases and dev) (push) Successful in 6m7s
release / Attach APK to the Release (tag releases only) (push) Skipped
release / Build + push container image (push) Successful in 1m36s
release / Verify release artifacts (tag releases only) (push) Skipped
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 <noreply@anthropic.com>
2026-10-08 09:24:20 -04:00

101 lines
3.4 KiB
Go

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)
}
}
}