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