Skip to main content

koan_core/
radio.rs

1//! Radio mode: picks tracks from the local library by several similarity axes:
2//! - ListenBrainz similar artists (no API key)
3//! - MusicBrainz relationships (collaborators, band members, associated acts)
4//! - Subsonic getSimilarSongs2 (when remote is configured)
5//! - Genre/era matching from local metadata
6//! - Acoustic similarity from stored feature vectors
7//! - Play history, which favours tracks not played recently
8//!
9//! The seed *drifts* — recent plays are weighted more heavily than the initial track,
10//! so the radio evolves through your library instead of orbiting one point.
11
12use std::collections::{HashMap, HashSet};
13use std::time::{SystemTime, UNIX_EPOCH};
14
15use rusqlite::Connection;
16
17use crate::config::RadioConfig;
18use crate::db::queries;
19use crate::remote::client::SubsonicClient;
20use crate::remote::listenbrainz;
21use crate::remote::musicbrainz;
22
23/// Which similarity axis led to a candidate being selected.
24#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
25pub enum SimilarityAxis {
26    ListenBrainz,
27    MusicBrainz,
28    Subsonic,
29    GenreEra,
30    SameArtist,
31    Random,
32    Acoustic,
33}
34
35/// A candidate track with its scoring breakdown.
36#[derive(Debug)]
37struct Candidate {
38    track_id: i64,
39    path: Option<String>,
40    year: Option<i32>,
41    /// Similarity axes that contributed to this candidate.
42    axes: HashSet<SimilarityAxis>,
43    /// Base similarity score (0.0..1.0).
44    base_score: f64,
45}
46
47/// Context extracted from the current queue and play history to guide radio picks.
48#[derive(Debug, Default)]
49pub struct RadioContext {
50    /// Whether the slow, network-backed signals may run. False when the queue
51    /// is about to run dry and a pick is needed this instant.
52    pub allow_network: bool,
53    /// Artist IDs from the seed window, with recency weight (more recent = higher).
54    pub seed_artists: HashMap<i64, f64>,
55    /// Paths already in the queue (to avoid duplicates).
56    pub queued_paths: HashSet<String>,
57    /// Track IDs in the recent play history exclusion window.
58    pub excluded_track_ids: HashSet<i64>,
59    /// The currently playing track's remote_id (for Subsonic similar songs).
60    pub current_remote_id: Option<String>,
61    /// The currently playing track's artist name (for top songs fallback).
62    pub current_artist_name: Option<String>,
63    /// Genres from the seed window.
64    pub seed_genres: HashSet<String>,
65    /// Average year of seed tracks (for era matching).
66    pub seed_avg_year: Option<i32>,
67}
68
69impl RadioContext {
70    /// Build context from queue items and play history.
71    ///
72    /// `queue_items`: (artist_id, path) pairs from the current queue.
73    /// `seed_window`: number of recent tracks to use as seeds.
74    /// `history_window`: number of recent track IDs to exclude.
75    pub fn build(
76        conn: &Connection,
77        queue_items: &[(Option<i64>, Option<String>)],
78        seed_window: usize,
79        history_window: usize,
80    ) -> Self {
81        let mut ctx = Self {
82            // Enrichment on by default; the caller turns it off when it cannot
83            // wait for it.
84            allow_network: true,
85            ..Self::default()
86        };
87
88        // Add queued paths for duplicate prevention.
89        for (_aid, path) in queue_items {
90            if let Some(p) = path {
91                ctx.queued_paths.insert(p.clone());
92            }
93        }
94
95        // Build seed from recent plays (drifting seed).
96        let recent =
97            queries::recent_track_ids(conn, queries::LOCAL_USER, seed_window).unwrap_or_default();
98        let seed_count = recent.len().max(1) as f64;
99        let mut years: Vec<i32> = Vec::new();
100
101        for (i, track_id) in recent.iter().enumerate() {
102            if let Ok(Some(track)) = queries::get_track_row(conn, *track_id) {
103                // More recent = higher weight (linear decay).
104                let weight = (seed_count - i as f64) / seed_count;
105                if let Some(aid) = track.artist_id {
106                    let entry = ctx.seed_artists.entry(aid).or_insert(0.0);
107                    *entry = entry.max(weight);
108                }
109                if let Some(ref genre) = track.genre {
110                    ctx.seed_genres.insert(genre.clone());
111                }
112                if let Some(album_id) = track.album_id
113                    && let Ok(Some(album)) = queries::get_album(conn, album_id)
114                    && let Some(ref date) = album.date
115                    && let Some(Ok(year)) = crate::helpers::year_of(date).map(str::parse::<i32>)
116                {
117                    years.push(year);
118                }
119            }
120        }
121
122        // With no play history, weight the artists in the queue instead.
123        if ctx.seed_artists.is_empty() {
124            for (artist_id, _path) in queue_items {
125                if let Some(aid) = artist_id {
126                    *ctx.seed_artists.entry(*aid).or_default() += 1.0;
127                }
128            }
129            // Normalise.
130            let max = ctx.seed_artists.values().copied().fold(1.0_f64, f64::max);
131            for v in ctx.seed_artists.values_mut() {
132                *v /= max;
133            }
134
135            // Collect genres from queue.
136            for (artist_id, _path) in queue_items {
137                if let Some(aid) = artist_id
138                    && let Ok(tracks) = queries::random_tracks_excluding(conn, &[], &[*aid], &[], 1)
139                {
140                    for t in tracks {
141                        if let Some(ref g) = t.genre {
142                            ctx.seed_genres.insert(g.clone());
143                        }
144                    }
145                }
146            }
147        }
148
149        // Build exclusion window from play history.
150        let excluded = queries::recent_track_ids(conn, queries::LOCAL_USER, history_window)
151            .unwrap_or_default();
152        ctx.excluded_track_ids = excluded.into_iter().collect();
153
154        // Average year of the seed tracks.
155        if !years.is_empty() {
156            ctx.seed_avg_year = Some(years.iter().sum::<i32>() / years.len() as i32);
157        }
158
159        ctx
160    }
161
162    /// Context from queue items alone, with no database behind it.
163    #[cfg(test)]
164    pub fn from_queue(items: &[(Option<i64>, Option<String>)]) -> Self {
165        let mut ctx = Self::default();
166        for (artist_id, path) in items {
167            if let Some(aid) = artist_id {
168                *ctx.seed_artists.entry(*aid).or_default() += 1.0;
169            }
170            if let Some(p) = path {
171                ctx.queued_paths.insert(p.clone());
172            }
173        }
174        // Normalise.
175        let max = ctx.seed_artists.values().copied().fold(1.0_f64, f64::max);
176        if max > 0.0 {
177            for v in ctx.seed_artists.values_mut() {
178                *v /= max;
179            }
180        }
181        ctx
182    }
183}
184
185/// Pick tracks for radio mode. Returns track IDs to enqueue.
186///
187/// Multi-signal strategy with fallback chain:
188/// 1. ListenBrainz similar artists -> local tracks
189/// 2. MusicBrainz relationships -> local tracks by collaborators/associated acts
190/// 3. Subsonic getSimilarSongs2 (if remote configured)
191/// 4. Genre + era match -> local tracks with matching tags from similar decade
192/// 5. Same-artist fallback
193/// 6. Acoustic similarity (vector KNN)
194/// 7. Random from library, the last resort
195pub fn pick_tracks(
196    conn: &Connection,
197    ctx: &RadioContext,
198    client: Option<&SubsonicClient>,
199    config: &RadioConfig,
200) -> Vec<i64> {
201    let count = config.batch_size;
202    let mut candidates: Vec<Candidate> = Vec::new();
203
204    log::info!(
205        "radio: picking {} tracks (seed: {} artists, {} genres, {} excluded, remote_id={}, artist={})",
206        count,
207        ctx.seed_artists.len(),
208        ctx.seed_genres.len(),
209        ctx.excluded_track_ids.len(),
210        ctx.current_remote_id.as_deref().unwrap_or("none"),
211        ctx.current_artist_name.as_deref().unwrap_or("none"),
212    );
213
214    // Signals 1-3 go to the network, and MusicBrainz is held to one request a
215    // second per seed artist, so they run only for a caller that can wait for
216    // them. The local signals are a database read.
217    if ctx.allow_network {
218        // --- Signal 1: ListenBrainz similar artists ---
219        gather_listenbrainz_candidates(conn, ctx, &mut candidates);
220
221        // --- Signal 2: MusicBrainz relationships ---
222        gather_musicbrainz_candidates(conn, ctx, &mut candidates);
223
224        // --- Signal 3: Subsonic similar songs ---
225        if let Some(client) = client {
226            gather_subsonic_candidates(conn, ctx, client, &mut candidates);
227        }
228    }
229
230    // --- Signal 4: Genre + era match ---
231    gather_genre_era_candidates(conn, ctx, &mut candidates);
232
233    // --- Signal 5: Same-artist tracks ---
234    gather_same_artist_candidates(conn, ctx, &mut candidates);
235
236    // --- Signal 6: Acoustic similarity (vector KNN) ---
237    gather_acoustic_candidates(conn, config.seed_window, &mut candidates);
238
239    // --- Signal 7: Random library tracks ---
240    gather_random_candidates(conn, ctx, &mut candidates);
241
242    log::info!("radio: {} raw candidates before scoring", candidates.len());
243
244    // Deduplicate by track_id, merging axes.
245    let mut deduped: HashMap<i64, Candidate> = HashMap::new();
246    for c in candidates {
247        let entry = deduped.entry(c.track_id).or_insert_with(|| Candidate {
248            track_id: c.track_id,
249            path: c.path.clone(),
250            year: c.year,
251            axes: HashSet::new(),
252            base_score: 0.0,
253        });
254        entry.axes.extend(c.axes.iter());
255        entry.base_score = entry.base_score.max(c.base_score);
256    }
257
258    // Filter out excluded tracks and already-queued.
259    let mut scored: Vec<(i64, f64)> = deduped
260        .into_values()
261        .filter(|c| !ctx.excluded_track_ids.contains(&c.track_id))
262        .filter(|c| {
263            c.path
264                .as_ref()
265                .is_none_or(|p| !ctx.queued_paths.contains(p))
266        })
267        .map(|c| {
268            let score = compute_score(conn, &c, ctx, config);
269            (c.track_id, score)
270        })
271        .collect();
272
273    // Sort by score descending.
274    scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
275
276    // Weighted random selection from top candidates for variety.
277    let picks = weighted_select(&scored, count);
278
279    log::info!("radio: picked {} tracks", picks.len());
280    picks
281}
282
283/// Compute final score for a candidate.
284fn compute_score(
285    conn: &Connection,
286    candidate: &Candidate,
287    ctx: &RadioContext,
288    config: &RadioConfig,
289) -> f64 {
290    let base = candidate.base_score;
291
292    // Signal overlap bonus: tracks matching on 2+ axes score higher.
293    let overlap_bonus = match candidate.axes.len() {
294        0 | 1 => 1.0,
295        2 => 1.5,
296        3 => 2.0,
297        _ => 2.5,
298    };
299
300    // Recency bonus: boost tracks that haven't been played recently or ever.
301    let recency_bonus = compute_recency_bonus(conn, candidate.track_id, config.discovery_weight);
302
303    // Era proximity bonus (if we have year data).
304    let era_bonus = if let (Some(track_year), Some(seed_year)) = (candidate.year, ctx.seed_avg_year)
305    {
306        let diff = (track_year - seed_year).unsigned_abs();
307        if diff <= 5 {
308            1.3
309        } else if diff <= 10 {
310            1.1
311        } else {
312            1.0
313        }
314    } else {
315        1.0
316    };
317
318    base * overlap_bonus * recency_bonus * era_bonus
319}
320
321/// Compute recency bonus for a track. Higher = more desirable.
322/// Never-played tracks get the highest bonus.
323fn compute_recency_bonus(conn: &Connection, track_id: i64, discovery_weight: f64) -> f64 {
324    let last_played = queries::last_played_at(conn, queries::LOCAL_USER, track_id).unwrap_or(None);
325    match last_played {
326        None => {
327            // Never played — big bonus, scaled by discovery_weight.
328            1.0 + discovery_weight * 2.0
329        }
330        Some(ts) => {
331            let now = SystemTime::now()
332                .duration_since(UNIX_EPOCH)
333                .unwrap_or_default()
334                .as_secs() as i64;
335            let days_ago = (now - ts) / 86400;
336            if days_ago > 180 {
337                1.0 + discovery_weight * 1.5 // Not heard in six months.
338            } else if days_ago > 30 {
339                1.0 + discovery_weight * 0.8
340            } else if days_ago > 7 {
341                1.0 + discovery_weight * 0.3
342            } else {
343                1.0 // Recently played — no bonus.
344            }
345        }
346    }
347}
348
349/// Weighted random selection from scored candidates.
350/// Takes the top N*3 candidates and selects N with probability proportional to score.
351fn weighted_select(scored: &[(i64, f64)], count: usize) -> Vec<i64> {
352    if scored.is_empty() {
353        return vec![];
354    }
355
356    let pool_size = (count * 3).min(scored.len());
357    let pool = &scored[..pool_size];
358
359    // Simple selection from the top-scored pool.
360    let mut selected = Vec::new();
361    let mut used = HashSet::new();
362
363    for (id, _score) in pool {
364        if selected.len() >= count {
365            break;
366        }
367        if used.insert(*id) {
368            selected.push(*id);
369        }
370    }
371
372    selected
373}
374
375// --- Signal gatherers ---
376
377fn gather_listenbrainz_candidates(
378    conn: &Connection,
379    ctx: &RadioContext,
380    candidates: &mut Vec<Candidate>,
381) {
382    let http = reqwest::blocking::Client::new();
383
384    for (&artist_id, &weight) in ctx.seed_artists.iter().take(3) {
385        // Get artist MBID from our DB.
386        let mbid: Option<String> = conn
387            .query_row(
388                "SELECT mbid FROM artists WHERE id = ?1",
389                rusqlite::params![artist_id],
390                |row| row.get(0),
391            )
392            .ok()
393            .flatten();
394
395        let mbid = match mbid {
396            Some(m) if !m.is_empty() => m,
397            _ => {
398                // Try to look up MBID via MusicBrainz search.
399                let artist_name: Option<String> = conn
400                    .query_row(
401                        "SELECT name FROM artists WHERE id = ?1",
402                        rusqlite::params![artist_id],
403                        |row| row.get(0),
404                    )
405                    .ok();
406                if let Some(name) = artist_name {
407                    match musicbrainz::lookup_artist_mbid(&http, &name) {
408                        Ok(Some(mbid)) => {
409                            // Cache the MBID.
410                            let _ = conn.execute(
411                                "UPDATE artists SET mbid = ?1 WHERE id = ?2",
412                                rusqlite::params![mbid, artist_id],
413                            );
414                            mbid
415                        }
416                        _ => continue,
417                    }
418                } else {
419                    continue;
420                }
421            }
422        };
423
424        // Check if we have fresh ListenBrainz data cached.
425        if queries::has_fresh_similar_artists_for_source(conn, artist_id, Some("listenbrainz"))
426            .unwrap_or(false)
427        {
428            // Use cached data.
429            add_cached_similar_candidates(
430                conn,
431                ctx,
432                artist_id,
433                weight,
434                SimilarityAxis::ListenBrainz,
435                candidates,
436            );
437            continue;
438        }
439
440        // Fetch from API.
441        match listenbrainz::get_similar_artists(&http, &mbid, 20) {
442            Ok(similar) => {
443                log::info!(
444                    "radio: listenbrainz returned {} similar for artist_id={}",
445                    similar.len(),
446                    artist_id
447                );
448                // Match to local artists and cache.
449                let mut pairs: Vec<(i64, f64)> = Vec::new();
450                for sa in &similar {
451                    // Try to find by MBID first, then by name.
452                    let local_id: Option<i64> = conn
453                        .query_row(
454                            "SELECT id FROM artists WHERE mbid = ?1",
455                            rusqlite::params![sa.mbid],
456                            |row| row.get(0),
457                        )
458                        .ok()
459                        .or_else(|| {
460                            conn.query_row(
461                                "SELECT id FROM artists WHERE name = ?1 COLLATE NOCASE",
462                                rusqlite::params![sa.name],
463                                |row| row.get(0),
464                            )
465                            .ok()
466                        });
467
468                    if let Some(local_id) = local_id
469                        && local_id != artist_id
470                    {
471                        pairs.push((local_id, sa.score));
472                    }
473                }
474
475                if !pairs.is_empty() {
476                    let _ = queries::save_similar_artists(conn, artist_id, &pairs, "listenbrainz");
477                }
478
479                // Add candidates from the matched local artists.
480                add_local_artist_candidates(
481                    conn,
482                    ctx,
483                    &pairs,
484                    weight,
485                    SimilarityAxis::ListenBrainz,
486                    candidates,
487                );
488            }
489            Err(e) => {
490                log::debug!(
491                    "radio: listenbrainz failed for artist_id={}: {}",
492                    artist_id,
493                    e
494                );
495                // Fall through — other signals will pick up the slack.
496            }
497        }
498    }
499}
500
501fn gather_musicbrainz_candidates(
502    conn: &Connection,
503    ctx: &RadioContext,
504    candidates: &mut Vec<Candidate>,
505) {
506    let http = musicbrainz::default_client();
507
508    for (&artist_id, &weight) in ctx.seed_artists.iter().take(3) {
509        // Check cache first.
510        if queries::has_fresh_similar_artists_for_source(conn, artist_id, Some("musicbrainz"))
511            .unwrap_or(false)
512        {
513            add_cached_similar_candidates(
514                conn,
515                ctx,
516                artist_id,
517                weight,
518                SimilarityAxis::MusicBrainz,
519                candidates,
520            );
521            continue;
522        }
523
524        let mbid: Option<String> = conn
525            .query_row(
526                "SELECT mbid FROM artists WHERE id = ?1",
527                rusqlite::params![artist_id],
528                |row| row.get(0),
529            )
530            .ok()
531            .flatten();
532
533        let Some(mbid) = mbid.filter(|m| !m.is_empty()) else {
534            continue;
535        };
536
537        match musicbrainz::get_artist_relations(&http, &mbid) {
538            Ok(relations) => {
539                log::info!(
540                    "radio: musicbrainz returned {} relations for artist_id={}",
541                    relations.len(),
542                    artist_id
543                );
544
545                let mut pairs: Vec<(i64, f64)> = Vec::new();
546
547                for rel in &relations {
548                    let local_id: Option<i64> = conn
549                        .query_row(
550                            "SELECT id FROM artists WHERE mbid = ?1",
551                            rusqlite::params![rel.mbid],
552                            |row| row.get(0),
553                        )
554                        .ok()
555                        .or_else(|| {
556                            conn.query_row(
557                                "SELECT id FROM artists WHERE name = ?1 COLLATE NOCASE",
558                                rusqlite::params![rel.name],
559                                |row| row.get(0),
560                            )
561                            .ok()
562                        });
563
564                    if let Some(local_id) = local_id
565                        && local_id != artist_id
566                    {
567                        // Score by relationship type.
568                        let score = match rel.category {
569                            musicbrainz::RelationCategory::Member => 0.8,
570                            musicbrainz::RelationCategory::Collaborator => 0.7,
571                            musicbrainz::RelationCategory::Associated => 0.5,
572                        };
573                        pairs.push((local_id, score));
574                    }
575                }
576
577                if !pairs.is_empty() {
578                    let _ = queries::save_similar_artists_with_rel(
579                        conn,
580                        artist_id,
581                        &pairs,
582                        "musicbrainz",
583                        "collaborator",
584                    );
585                }
586
587                add_local_artist_candidates(
588                    conn,
589                    ctx,
590                    &pairs,
591                    weight,
592                    SimilarityAxis::MusicBrainz,
593                    candidates,
594                );
595            }
596            Err(e) => {
597                log::debug!(
598                    "radio: musicbrainz relations failed for artist_id={}: {}",
599                    artist_id,
600                    e
601                );
602            }
603        }
604
605        // Rate limit: sleep 1s between MusicBrainz requests.
606        std::thread::sleep(std::time::Duration::from_secs(1));
607    }
608}
609
610fn gather_subsonic_candidates(
611    conn: &Connection,
612    ctx: &RadioContext,
613    client: &SubsonicClient,
614    candidates: &mut Vec<Candidate>,
615) {
616    if let Some(ref remote_id) = ctx.current_remote_id {
617        match client.get_similar_songs(remote_id, 30) {
618            Ok(songs) => {
619                log::info!("radio: subsonic returned {} similar songs", songs.len());
620                for (i, song) in songs.iter().enumerate() {
621                    if let Some(track_id) = resolve_subsonic_song_to_track(conn, song) {
622                        let score = (songs.len() as f64 - i as f64) / songs.len() as f64;
623                        let track = queries::get_track_row(conn, track_id).ok().flatten();
624                        candidates.push(Candidate {
625                            track_id,
626                            path: track.as_ref().and_then(|t| t.path.clone()),
627                            year: None,
628                            axes: [SimilarityAxis::Subsonic].into_iter().collect(),
629                            base_score: score * 0.9,
630                        });
631                    }
632                }
633
634                // Cache artist relationships from subsonic results.
635                cache_subsonic_artist_relationships(conn, ctx, &songs);
636            }
637            Err(e) => {
638                log::debug!("radio: subsonic similar songs failed: {}", e);
639            }
640        }
641    }
642}
643
644fn gather_genre_era_candidates(
645    conn: &Connection,
646    ctx: &RadioContext,
647    candidates: &mut Vec<Candidate>,
648) {
649    if ctx.seed_genres.is_empty() {
650        return;
651    }
652
653    let genres: Vec<String> = ctx.seed_genres.iter().cloned().collect();
654    let exclude: Vec<String> = ctx.queued_paths.iter().cloned().collect();
655
656    match queries::random_tracks_excluding(conn, &exclude, &[], &genres, 30) {
657        Ok(tracks) => {
658            for track in tracks {
659                let year = track
660                    .album_id
661                    .and_then(|aid| queries::get_album(conn, aid).ok().flatten())
662                    .and_then(|a| {
663                        a.date
664                            .as_ref()
665                            .and_then(|d| crate::helpers::year_of(d)?.parse().ok())
666                    });
667
668                // Score higher if both genre AND era match.
669                let genre_match = track
670                    .genre
671                    .as_ref()
672                    .is_some_and(|g| ctx.seed_genres.contains(g));
673                let era_match = match (year, ctx.seed_avg_year) {
674                    (Some(y), Some(sy)) => (y as i64 - sy as i64).unsigned_abs() <= 10,
675                    _ => false,
676                };
677
678                let base_score = match (genre_match, era_match) {
679                    (true, true) => 0.6,
680                    (true, false) => 0.3,
681                    (false, true) => 0.2,
682                    (false, false) => 0.1,
683                };
684
685                candidates.push(Candidate {
686                    track_id: track.id,
687                    path: track.path.clone(),
688                    year,
689                    axes: [SimilarityAxis::GenreEra].into_iter().collect(),
690                    base_score,
691                });
692            }
693        }
694        Err(e) => {
695            log::debug!("radio: genre/era query failed: {}", e);
696        }
697    }
698}
699
700fn gather_same_artist_candidates(
701    conn: &Connection,
702    ctx: &RadioContext,
703    candidates: &mut Vec<Candidate>,
704) {
705    let artist_ids: Vec<i64> = ctx.seed_artists.keys().copied().collect();
706    if artist_ids.is_empty() {
707        return;
708    }
709
710    let exclude: Vec<String> = ctx.queued_paths.iter().cloned().collect();
711    match queries::random_tracks_excluding(conn, &exclude, &artist_ids, &[], 15) {
712        Ok(tracks) => {
713            for track in tracks {
714                let weight = track
715                    .artist_id
716                    .and_then(|aid| ctx.seed_artists.get(&aid))
717                    .copied()
718                    .unwrap_or(0.3);
719
720                candidates.push(Candidate {
721                    track_id: track.id,
722                    path: track.path.clone(),
723                    year: None,
724                    axes: [SimilarityAxis::SameArtist].into_iter().collect(),
725                    base_score: weight * 0.4, // Lower base — same-artist is the fallback.
726                });
727            }
728        }
729        Err(e) => {
730            log::debug!("radio: same-artist query failed: {}", e);
731        }
732    }
733}
734
735fn gather_acoustic_candidates(
736    conn: &Connection,
737    seed_window: usize,
738    candidates: &mut Vec<Candidate>,
739) {
740    // The same seeds that drive seed_artists — one window, one answer.
741    let seed_ids =
742        queries::recent_track_ids(conn, queries::LOCAL_USER, seed_window).unwrap_or_default();
743    let mut seed_embeddings = Vec::new();
744    for tid in &seed_ids {
745        if let Ok(Some(emb)) = queries::get_vector(conn, *tid) {
746            seed_embeddings.push(emb);
747        }
748    }
749
750    if seed_embeddings.is_empty() {
751        return;
752    }
753
754    let centroid = crate::index::features::centroid(&seed_embeddings);
755    let knn_result = queries::find_similar_to_vector(conn, &centroid, 30, None);
756    match knn_result {
757        Ok(ref results) => {
758            let max_dist = results.last().map(|r| r.1).unwrap_or(1.0).max(0.001);
759            let mut added = 0;
760            for &(track_id, dist) in results {
761                // Skip seed tracks themselves.
762                if seed_ids.contains(&track_id) {
763                    continue;
764                }
765                // Score: inverse of normalised distance. Closer = higher score.
766                let score = (1.0 - (dist / max_dist)).max(0.0) as f64 * 0.7;
767                let track = queries::get_track_row(conn, track_id).ok().flatten();
768                candidates.push(Candidate {
769                    track_id,
770                    path: track.as_ref().and_then(|t| t.path.clone()),
771                    year: None,
772                    axes: [SimilarityAxis::Acoustic].into_iter().collect(),
773                    base_score: score,
774                });
775                added += 1;
776            }
777            log::info!(
778                "radio: acoustic signal added {} candidates from {} seed vectors",
779                added,
780                seed_embeddings.len()
781            );
782        }
783        Err(e) => {
784            log::debug!("radio: acoustic similarity query failed: {}", e);
785        }
786    }
787}
788
789fn gather_random_candidates(
790    conn: &Connection,
791    ctx: &RadioContext,
792    candidates: &mut Vec<Candidate>,
793) {
794    let exclude: Vec<String> = ctx.queued_paths.iter().cloned().collect();
795    match queries::random_tracks_excluding(conn, &exclude, &[], &[], 10) {
796        Ok(tracks) => {
797            for track in tracks {
798                candidates.push(Candidate {
799                    track_id: track.id,
800                    path: track.path.clone(),
801                    year: None,
802                    axes: [SimilarityAxis::Random].into_iter().collect(),
803                    base_score: 0.05, // Last resort: any track beats silence.
804                });
805            }
806        }
807        Err(e) => {
808            log::debug!("radio: random fallback failed: {}", e);
809        }
810    }
811}
812
813// --- Helpers ---
814
815/// Add candidates from cached similar artist data.
816fn add_cached_similar_candidates(
817    conn: &Connection,
818    ctx: &RadioContext,
819    artist_id: i64,
820    seed_weight: f64,
821    axis: SimilarityAxis,
822    candidates: &mut Vec<Candidate>,
823) {
824    if let Ok(similar) = queries::get_similar_artists(conn, artist_id) {
825        let pairs: Vec<(i64, f64)> = similar.into_iter().map(|(a, s)| (a.id, s)).collect();
826        add_local_artist_candidates(conn, ctx, &pairs, seed_weight, axis, candidates);
827    }
828}
829
830/// Add candidates from a list of (artist_id, similarity_score) pairs.
831fn add_local_artist_candidates(
832    conn: &Connection,
833    ctx: &RadioContext,
834    pairs: &[(i64, f64)],
835    seed_weight: f64,
836    axis: SimilarityAxis,
837    candidates: &mut Vec<Candidate>,
838) {
839    let exclude: Vec<String> = ctx.queued_paths.iter().cloned().collect();
840    for &(similar_artist_id, sim_score) in pairs.iter().take(10) {
841        if let Ok(tracks) =
842            queries::random_tracks_excluding(conn, &exclude, &[similar_artist_id], &[], 3)
843        {
844            for track in tracks {
845                candidates.push(Candidate {
846                    track_id: track.id,
847                    path: track.path.clone(),
848                    year: None,
849                    axes: [axis].into_iter().collect(),
850                    base_score: sim_score * seed_weight * 0.8,
851                });
852            }
853        }
854    }
855}
856
857/// Resolve a SubsonicSong to a local track ID by remote_id.
858fn resolve_subsonic_song_to_track(
859    conn: &Connection,
860    song: &crate::remote::client::SubsonicSong,
861) -> Option<i64> {
862    conn.query_row(
863        "SELECT id FROM tracks WHERE remote_id = ?1",
864        rusqlite::params![song.id],
865        |row| row.get::<_, i64>(0),
866    )
867    .ok()
868}
869
870/// Extract and cache artist relationships from Subsonic similar songs response.
871fn cache_subsonic_artist_relationships(
872    conn: &Connection,
873    ctx: &RadioContext,
874    songs: &[crate::remote::client::SubsonicSong],
875) {
876    for &artist_id in ctx.seed_artists.keys().take(5) {
877        let mut similar: HashMap<i64, f64> = HashMap::new();
878        let total = songs.len() as f64;
879
880        for (i, song) in songs.iter().enumerate() {
881            if let Some(ref song_artist_id) = song.artist_id {
882                let local_artist_id: Option<i64> = conn
883                    .query_row(
884                        "SELECT id FROM artists WHERE remote_id = ?1",
885                        rusqlite::params![song_artist_id],
886                        |row| row.get(0),
887                    )
888                    .ok();
889
890                if let Some(local_id) = local_artist_id
891                    && local_id != artist_id
892                {
893                    let score = (total - i as f64) / total;
894                    let entry = similar.entry(local_id).or_insert(0.0);
895                    *entry = entry.max(score);
896                }
897            }
898        }
899
900        if !similar.is_empty() {
901            let pairs: Vec<(i64, f64)> = similar.into_iter().collect();
902            let _ = queries::save_similar_artists(conn, artist_id, &pairs, "subsonic");
903        }
904    }
905}
906
907// ---------------------------------------------------------------------------
908// Auto-queue
909// ---------------------------------------------------------------------------
910
911/// Keep the queue topped up while radio mode is on.
912///
913/// Radio mode is a flag on `SharedPlayerState`. Owning the loop here means
914/// every front end that sets it gets the same behaviour instead of
915/// reimplementing it.
916///
917/// Runs on its own thread and exits when the player goes away.
918pub fn spawn_autoqueue(
919    state: std::sync::Arc<crate::player::state::SharedPlayerState>,
920    tx: crossbeam_channel::Sender<crate::player::commands::PlayerCommand>,
921    db_path: std::path::PathBuf,
922) {
923    use crate::player::commands::PlayerCommand;
924    use crate::player::state::{QueueEntryStatus, QueueItemId};
925
926    std::thread::Builder::new()
927        .name("koan-radio".into())
928        .spawn(move || {
929            let pool = crate::db::pool::Pool::new(db_path);
930            loop {
931                std::thread::sleep(std::time::Duration::from_secs(2));
932
933                if !state.radio_mode() || state.cursor().is_none() {
934                    continue;
935                }
936                log::debug!("radio: awake, cursor set");
937
938                let cfg = crate::config::Config::load().unwrap_or_default();
939                let snapshot = state.derive_visible_queue();
940                let Some(playing) = snapshot
941                    .entries
942                    .iter()
943                    .position(|e| e.status == QueueEntryStatus::Playing)
944                else {
945                    log::debug!("radio: nothing is playing, waiting");
946                    continue;
947                };
948                let remaining = snapshot
949                    .entries
950                    .iter()
951                    .skip(playing + 1)
952                    .filter(|e| e.status == QueueEntryStatus::Queued)
953                    .count();
954                if remaining > cfg.radio.lookahead {
955                    continue;
956                }
957                log::info!(
958                    "radio: {} queued after the cursor, topping up to {}",
959                    remaining,
960                    cfg.radio.lookahead
961                );
962
963                let Ok(db) = pool.get() else {
964                    continue;
965                };
966                let (items, cursor) = state.snapshot_playlist();
967
968                // The seed drifts: recent items weigh more than the first thing
969                // queued, so the radio moves through the library rather than
970                // orbiting one track.
971                let context: Vec<(Option<i64>, Option<String>)> = items
972                    .iter()
973                    .map(|item| {
974                        let row = item
975                            .path
976                            .to_str()
977                            .and_then(|p| queries::track_id_by_path(&db.conn, p).ok())
978                            .flatten()
979                            .and_then(|id| queries::get_track_row(&db.conn, id).ok())
980                            .flatten();
981                        (
982                            row.as_ref().and_then(|t| t.artist_id),
983                            Some(item.path.to_string_lossy().into_owned()),
984                        )
985                    })
986                    .collect();
987
988                let mut ctx = RadioContext::build(
989                    &db.conn,
990                    &context,
991                    cfg.radio.seed_window,
992                    cfg.radio.history_window,
993                );
994                if let Some(current) = cursor.and_then(|cid| items.iter().find(|i| i.id == cid))
995                    && let Some(row) = current
996                        .path
997                        .to_str()
998                        .and_then(|p| queries::track_id_by_path(&db.conn, p).ok())
999                        .flatten()
1000                        .and_then(|id| queries::get_track_row(&db.conn, id).ok())
1001                        .flatten()
1002                {
1003                    ctx.current_remote_id = row.remote_id.clone();
1004                    ctx.current_artist_name = Some(row.artist_name.clone());
1005                }
1006                // Local signals only.
1007                //
1008                // ListenBrainz and MusicBrainz each rate-limit to one request a
1009                // second per seed artist, and both are called in line before a
1010                // single pick comes back — so a queue that needs a track in the
1011                // next few seconds gets one long after the music has stopped.
1012                // Genre/era, same-artist, acoustic similarity and plain random
1013                // are all database reads and answer immediately, which is worth
1014                // more than a better-chosen track that arrives too late.
1015
1016                ctx.allow_network = false;
1017
1018                // No client, and no similar-artist prefetch: both are HTTP
1019                // round trips in front of a pick that is needed now. Cached
1020                // similar artists are still read from the database by the local
1021                // signals; nothing is fetched.
1022                let picks = pick_tracks(&db.conn, &ctx, None, &cfg.radio);
1023                if picks.is_empty() {
1024                    log::warn!("radio: the picker returned nothing for this seed");
1025                    continue;
1026                }
1027
1028                // Never queue something already in the queue: the picker scores
1029                // by similarity and has no idea what is sitting below the
1030                // cursor.
1031                let queued: HashSet<String> = items
1032                    .iter()
1033                    .map(|i| i.path.to_string_lossy().into_owned())
1034                    .collect();
1035                let rows: Vec<_> = queries::tracks_by_ids(&db.conn, &picks)
1036                    .unwrap_or_default()
1037                    .into_iter()
1038                    .filter(|row| {
1039                        row.path
1040                            .as_deref()
1041                            .or(row.cached_path.as_deref())
1042                            .is_none_or(|p| !queued.contains(p))
1043                    })
1044                    .collect();
1045                if rows.is_empty() {
1046                    log::warn!("radio: every pick was already in the queue");
1047                    continue;
1048                }
1049
1050                let new_items = crate::helpers::playlist_items_for_tracks(&db, &rows);
1051                let pending: Vec<(i64, QueueItemId)> = new_items
1052                    .iter()
1053                    .filter(|i| matches!(i.state, crate::player::state::ItemState::Pending))
1054                    .filter_map(|i| i.db_id.map(|id| (id, i.id)))
1055                    .collect();
1056
1057                log::info!("radio: queueing {} tracks", new_items.len());
1058                if tx.send(PlayerCommand::AddToPlaylist(new_items)).is_err() {
1059                    return; // Player gone; so is the app.
1060                }
1061                if !pending.is_empty() {
1062                    crate::helpers::spawn_downloads(pending, tx.clone(), state.clone());
1063                }
1064            }
1065        })
1066        .ok();
1067}
1068
1069#[cfg(test)]
1070mod tests {
1071    use super::*;
1072    use crate::db::connection::Database;
1073    use crate::db::queries::{get_or_create_artist, sample_meta, upsert_track};
1074
1075    fn test_db() -> Database {
1076        let conn = rusqlite::Connection::open_in_memory().unwrap();
1077        conn.pragma_update(None, "foreign_keys", "on").unwrap();
1078        crate::db::schema::create_tables(&conn).unwrap();
1079        Database { conn }
1080    }
1081
1082    #[test]
1083    fn test_radio_context_from_queue() {
1084        let ctx = RadioContext::from_queue(&[
1085            (Some(1), Some("/a.flac".into())),
1086            (Some(2), Some("/b.flac".into())),
1087            (Some(1), Some("/c.flac".into())),
1088        ]);
1089        assert_eq!(ctx.seed_artists.len(), 2);
1090        assert!(ctx.seed_artists[&1] > ctx.seed_artists[&2]);
1091        assert_eq!(ctx.queued_paths.len(), 3);
1092    }
1093
1094    #[test]
1095    fn test_recency_bonus_never_played() {
1096        let db = test_db();
1097        let mut meta = sample_meta("T1", "A1", "Al1");
1098        meta.path = Some("/music/T1.flac".into());
1099        upsert_track(&db.conn, &meta).unwrap();
1100
1101        let track_id: i64 = db
1102            .conn
1103            .query_row("SELECT id FROM tracks LIMIT 1", [], |row| row.get(0))
1104            .unwrap();
1105
1106        let bonus = compute_recency_bonus(&db.conn, track_id, 0.3);
1107        assert!(bonus > 1.0, "never-played should get a bonus");
1108    }
1109
1110    #[test]
1111    fn test_recency_bonus_recently_played() {
1112        let db = test_db();
1113        let mut meta = sample_meta("T1", "A1", "Al1");
1114        meta.path = Some("/music/T1.flac".into());
1115        upsert_track(&db.conn, &meta).unwrap();
1116
1117        let track_id: i64 = db
1118            .conn
1119            .query_row("SELECT id FROM tracks LIMIT 1", [], |row| row.get(0))
1120            .unwrap();
1121
1122        queries::record_play(
1123            &db.conn,
1124            crate::db::queries::LOCAL_USER,
1125            track_id,
1126            Some(240_000),
1127        )
1128        .unwrap();
1129
1130        let bonus = compute_recency_bonus(&db.conn, track_id, 0.3);
1131        assert!(
1132            (bonus - 1.0).abs() < f64::EPSILON,
1133            "recently played should get no bonus"
1134        );
1135    }
1136
1137    #[test]
1138    fn test_weighted_select_empty() {
1139        assert!(weighted_select(&[], 5).is_empty());
1140    }
1141
1142    #[test]
1143    fn test_weighted_select_fewer_than_requested() {
1144        let scored = vec![(1, 0.9), (2, 0.5)];
1145        let picks = weighted_select(&scored, 5);
1146        assert_eq!(picks.len(), 2);
1147    }
1148
1149    #[test]
1150    fn test_signal_overlap_scoring() {
1151        let db = test_db();
1152        let config = RadioConfig::default();
1153        let ctx = RadioContext::default();
1154
1155        // Candidate with 1 axis.
1156        let c1 = Candidate {
1157            track_id: 1,
1158            path: None,
1159            year: None,
1160            axes: [SimilarityAxis::ListenBrainz].into_iter().collect(),
1161            base_score: 0.5,
1162        };
1163
1164        // Candidate with 3 axes.
1165        let c3 = Candidate {
1166            track_id: 2,
1167            path: None,
1168            year: None,
1169            axes: [
1170                SimilarityAxis::ListenBrainz,
1171                SimilarityAxis::MusicBrainz,
1172                SimilarityAxis::GenreEra,
1173            ]
1174            .into_iter()
1175            .collect(),
1176            base_score: 0.5,
1177        };
1178
1179        let score1 = compute_score(&db.conn, &c1, &ctx, &config);
1180        let score3 = compute_score(&db.conn, &c3, &ctx, &config);
1181
1182        assert!(
1183            score3 > score1,
1184            "multi-axis candidate should score higher: {} vs {}",
1185            score3,
1186            score1
1187        );
1188    }
1189
1190    #[test]
1191    fn test_pick_tracks_empty_library() {
1192        let db = test_db();
1193        let ctx = RadioContext::from_queue(&[]);
1194        let config = RadioConfig::default();
1195        let picks = pick_tracks(&db.conn, &ctx, None, &config);
1196        assert!(picks.is_empty());
1197    }
1198
1199    #[test]
1200    fn test_pick_tracks_with_library() {
1201        let db = test_db();
1202
1203        // Populate library.
1204        for i in 0..20 {
1205            let mut meta = sample_meta(
1206                &format!("Track{}", i),
1207                &format!("Artist{}", i % 5),
1208                &format!("Album{}", i % 3),
1209            );
1210            meta.path = Some(format!("/music/Album{}/Track{}.flac", i % 3, i));
1211            meta.track_number = Some(i);
1212            upsert_track(&db.conn, &meta).unwrap();
1213        }
1214
1215        let artist_id: i64 = db
1216            .conn
1217            .query_row("SELECT id FROM artists LIMIT 1", [], |row| row.get(0))
1218            .unwrap();
1219
1220        let ctx = RadioContext::from_queue(&[(Some(artist_id), Some("/queued.flac".into()))]);
1221        let config = RadioConfig {
1222            batch_size: 5,
1223            ..RadioConfig::default()
1224        };
1225
1226        let picks = pick_tracks(&db.conn, &ctx, None, &config);
1227        assert!(
1228            !picks.is_empty(),
1229            "should pick at least some tracks from a populated library"
1230        );
1231        assert!(picks.len() <= 5);
1232    }
1233
1234    #[test]
1235    fn test_pick_tracks_excludes_history() {
1236        let db = test_db();
1237
1238        // Insert a few tracks.
1239        for i in 0..5 {
1240            let mut meta = sample_meta(&format!("T{}", i), "Artist", "Album");
1241            meta.path = Some(format!("/music/T{}.flac", i));
1242            meta.track_number = Some(i);
1243            upsert_track(&db.conn, &meta).unwrap();
1244        }
1245
1246        // Record all as recently played.
1247        let mut ids = Vec::new();
1248        for i in 0..5 {
1249            let id: i64 = db
1250                .conn
1251                .query_row(
1252                    "SELECT id FROM tracks WHERE path = ?1",
1253                    rusqlite::params![format!("/music/T{}.flac", i)],
1254                    |row| row.get(0),
1255                )
1256                .unwrap();
1257            queries::record_play(&db.conn, crate::db::queries::LOCAL_USER, id, Some(240_000))
1258                .unwrap();
1259            ids.push(id);
1260        }
1261
1262        let mut ctx = RadioContext::from_queue(&[]);
1263        ctx.excluded_track_ids = ids.into_iter().collect();
1264
1265        let config = RadioConfig {
1266            batch_size: 5,
1267            ..RadioConfig::default()
1268        };
1269
1270        let picks = pick_tracks(&db.conn, &ctx, None, &config);
1271        // All tracks are excluded, so nothing should be picked.
1272        assert!(
1273            picks.is_empty(),
1274            "all tracks in exclusion window, got picks"
1275        );
1276    }
1277
1278    #[test]
1279    fn test_radio_context_build_with_no_history() {
1280        let db = test_db();
1281
1282        for i in 0..5 {
1283            let _id = get_or_create_artist(&db.conn, &format!("Artist{}", i), None).unwrap();
1284        }
1285
1286        let queue = vec![
1287            (Some(1_i64), Some("/a.flac".to_string())),
1288            (Some(2), Some("/b.flac".to_string())),
1289        ];
1290        let ctx = RadioContext::build(&db.conn, &queue, 5, 200);
1291
1292        // Should fall back to queue weights since no play history.
1293        assert_eq!(ctx.seed_artists.len(), 2);
1294        assert_eq!(ctx.queued_paths.len(), 2);
1295    }
1296}