/
/
/
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_similar_lookup_failure_keeps_base_tracks() -> None:
177 """When no provider supports similar tracks, base tracks still carry the playlist."""
178 seed = _seed(MediaType.TRACK, "s1")
179 prov = _make_provider({"s1": [_track("s1")]}, [])
180 cast("Any", prov.mass.music.tracks.similar_tracks).side_effect = UnsupportedFeaturedException(
181 "no similar provider"
182 )
183 result = await prov.get_dynamic_tracks([seed], include_base_tracks=True, target_size=5)
184 assert any(t.item_id == "s1" for t in result)
185