/
/
/
1"""Tests for the radio_playlist provider's track generation (base + similar assembly)."""
2
3from __future__ import annotations
4
5from typing import TYPE_CHECKING, Any, cast
6from unittest.mock import AsyncMock, MagicMock
7
8import pytest
9from music_assistant_models.enums import MediaType
10from music_assistant_models.errors import UnsupportedFeaturedException
11
12from music_assistant.controllers.music.constants import RADIO_TRACK_MAX_DURATION_SECS
13from music_assistant.providers.radio_playlist import RadioPlaylistProvider
14
15if TYPE_CHECKING:
16 from music_assistant import MusicAssistant
17
18
19def _track(item_id: str, duration: int = 200) -> MagicMock:
20 track = MagicMock()
21 track.item_id = item_id
22 track.provider = "test"
23 track.uri = f"test://track/{item_id}"
24 track.name = f"Track {item_id}"
25 track.duration = duration
26 track.media_type = MediaType.TRACK
27 return track
28
29
30def _seed(media_type: MediaType, item_id: str) -> MagicMock:
31 item = MagicMock()
32 item.media_type = media_type
33 item.item_id = item_id
34 item.provider = "test"
35 item.uri = f"test://{media_type.value}/{item_id}"
36 return item
37
38
39def _make_provider(
40 base_tracks_by_seed: dict[str, list[MagicMock]], similar: list[MagicMock]
41) -> RadioPlaylistProvider:
42 """Build a provider with the queue's track resolution + similar lookups stubbed out."""
43 prov = RadioPlaylistProvider.__new__(RadioPlaylistProvider)
44 mass = MagicMock()
45 mass.music.tracks.similar_tracks = AsyncMock(return_value=similar)
46 mass.player_queues.get_tracks_for_playback = AsyncMock(
47 side_effect=lambda item: base_tracks_by_seed.get(item.item_id, [])
48 )
49 prov.mass = cast("MusicAssistant", mass)
50 return prov
51
52
53@pytest.mark.asyncio
54async def test_no_seeds_raises() -> None:
55 """An empty seed list raises UnsupportedFeaturedException."""
56 prov = _make_provider({}, [])
57 with pytest.raises(UnsupportedFeaturedException):
58 await prov.get_dynamic_tracks([])
59
60
61@pytest.mark.asyncio
62async def test_seed_with_no_base_tracks_raises() -> None:
63 """When all seeds yield zero base tracks, raises UnsupportedFeaturedException."""
64 seed = _seed(MediaType.ALBUM, "1")
65 prov = _make_provider({"1": []}, [])
66 with pytest.raises(UnsupportedFeaturedException):
67 await prov.get_dynamic_tracks([seed])
68
69
70@pytest.mark.asyncio
71async def test_long_tracks_are_filtered() -> None:
72 """Tracks longer than the max-duration threshold are dropped from candidates."""
73 seed = _seed(MediaType.TRACK, "s1")
74 base = _track("s1")
75 short = _track("short", duration=100)
76 long_track = _track("long", duration=RADIO_TRACK_MAX_DURATION_SECS + 1)
77 prov = _make_provider({"s1": [base]}, [short, long_track])
78
79 result = await prov.get_dynamic_tracks([seed], target_size=10)
80 ids = {t.item_id for t in result}
81 assert "short" in ids
82 assert "long" not in ids
83
84
85@pytest.mark.asyncio
86async def test_include_base_tracks_emits_seed_track() -> None:
87 """include_base_tracks=True ensures at least one base track is in the output."""
88 seed = _seed(MediaType.TRACK, "s1")
89 base = _track("s1")
90 similar = [_track(f"sim{i}") for i in range(5)]
91 prov = _make_provider({"s1": [base]}, similar)
92
93 result = await prov.get_dynamic_tracks([seed], include_base_tracks=True, target_size=5)
94 assert any(t.item_id == "s1" for t in result)
95
96
97@pytest.mark.asyncio
98async def test_exclude_base_tracks_omits_seed_track() -> None:
99 """include_base_tracks=False keeps only similar tracks in the output."""
100 seed = _seed(MediaType.TRACK, "s1")
101 base = _track("s1")
102 similar = [_track(f"sim{i}") for i in range(5)]
103 prov = _make_provider({"s1": [base]}, similar)
104
105 result = await prov.get_dynamic_tracks([seed], include_base_tracks=False, target_size=5)
106 assert all(t.item_id != "s1" for t in result)
107 assert len(result) == 5
108
109
110@pytest.mark.asyncio
111async def test_multiple_seeds_dedup_base_tracks() -> None:
112 """Duplicate base tracks across seeds are deduplicated before sampling."""
113 seed_a = _seed(MediaType.ALBUM, "a")
114 seed_b = _seed(MediaType.ALBUM, "b")
115 shared = _track("shared")
116 prov = _make_provider(
117 {"a": [shared, _track("a-only")], "b": [shared, _track("b-only")]},
118 [],
119 )
120 similar_tracks = cast("Any", prov.mass.music.tracks.similar_tracks)
121 similar_tracks.side_effect = lambda item_id, _provider, **_kwargs: [
122 _track(f"similar-{item_id}")
123 ]
124
125 await prov.get_dynamic_tracks([seed_a, seed_b], target_size=5)
126 # Three unique base tracks remain after deduplication, each with a successful direct lookup.
127 assert similar_tracks.call_count == 3
128
129
130@pytest.mark.asyncio
131async def test_cross_provider_lookup_only_retries_without_usable_results() -> None:
132 """A cross-provider lookup is only used when the direct lookup has no usable result."""
133 seed = _seed(MediaType.ALBUM, "seed")
134 direct_base = _track("direct")
135 fallback_base = _track("fallback")
136 filtered_base = _track("filtered")
137 prov = _make_provider({"seed": [direct_base, fallback_base, filtered_base]}, [])
138 similar_tracks = cast("Any", prov.mass.music.tracks.similar_tracks)
139
140 def get_similar(
141 item_id: str, _provider: str, *, allow_lookup: bool, **_kwargs: Any
142 ) -> list[MagicMock]:
143 if item_id == "direct":
144 return [_track("direct-result")]
145 if item_id == "filtered":
146 return (
147 [_track("filtered-fallback")]
148 if allow_lookup
149 else [_track("too-long", RADIO_TRACK_MAX_DURATION_SECS + 1)]
150 )
151 if allow_lookup:
152 return [_track("fallback-result")]
153 return []
154
155 similar_tracks.side_effect = get_similar
156
157 result = await prov.get_dynamic_tracks([seed], target_size=5)
158
159 lookups = [
160 (call.args[0], call.kwargs["allow_lookup"]) for call in similar_tracks.await_args_list
161 ]
162 assert ("direct", False) in lookups
163 assert ("direct", True) not in lookups
164 assert ("fallback", False) in lookups
165 assert ("fallback", True) in lookups
166 assert ("filtered", False) in lookups
167 assert ("filtered", True) in lookups
168 assert {track.item_id for track in result} == {
169 "direct-result",
170 "fallback-result",
171 "filtered-fallback",
172 }
173
174
175@pytest.mark.asyncio
176async def test_artist_seed_uses_top_tracks_not_full_resolution() -> None:
177 """An artist seed samples base tracks via top_tracks; other seeds still go through the queue."""
178 artist_seed = _seed(MediaType.ARTIST, "art1")
179 album_seed = _seed(MediaType.ALBUM, "alb1")
180 prov = _make_provider({"alb1": [_track("albtrack1")]}, [])
181 prov.mass.music.artists.top_tracks = AsyncMock( # type: ignore[method-assign]
182 return_value=[_track("top1")]
183 )
184
185 await prov.get_dynamic_tracks([artist_seed, album_seed], target_size=5)
186
187 prov.mass.music.artists.top_tracks.assert_awaited_once_with("art1", "test")
188 get_tracks_for_playback = cast("Any", prov.mass.player_queues.get_tracks_for_playback)
189 called_item_ids = [call.args[0].item_id for call in get_tracks_for_playback.await_args_list]
190 assert called_item_ids == ["alb1"]
191
192
193@pytest.mark.asyncio
194async def test_artist_seed_without_top_tracks_falls_back() -> None:
195 """An artist whose providers yield no top tracks seeds via the plain tracks listing."""
196 artist_seed = _seed(MediaType.ARTIST, "art1")
197 prov = _make_provider({}, [_track("sim1")])
198 prov.mass.music.artists.top_tracks = AsyncMock(return_value=[]) # type: ignore[method-assign]
199 prov.mass.music.artists.tracks = AsyncMock( # type: ignore[method-assign]
200 return_value=[_track("full1")]
201 )
202
203 result = await prov.get_dynamic_tracks([artist_seed], include_base_tracks=True, target_size=5)
204
205 prov.mass.music.artists.tracks.assert_awaited_once_with("art1", "test")
206 cast("Any", prov.mass.player_queues.get_tracks_for_playback).assert_not_awaited()
207 assert any(t.item_id == "full1" for t in result)
208
209
210@pytest.mark.asyncio
211async def test_similar_lookup_failure_keeps_base_tracks() -> None:
212 """When no provider supports similar tracks, base tracks still carry the playlist."""
213 seed = _seed(MediaType.TRACK, "s1")
214 prov = _make_provider({"s1": [_track("s1")]}, [])
215 cast("Any", prov.mass.music.tracks.similar_tracks).side_effect = UnsupportedFeaturedException(
216 "no similar provider"
217 )
218 result = await prov.get_dynamic_tracks([seed], include_base_tracks=True, target_size=5)
219 assert any(t.item_id == "s1" for t in result)
220