/
/
/
1"""Tests for ProviderFeature.SEARCH wiring and the search() dispatcher hook."""
2
3from __future__ import annotations
4
5from typing import TYPE_CHECKING
6from unittest.mock import AsyncMock, MagicMock
7
8import numpy as np
9import pytest
10from music_assistant_models.enums import MediaType, ProviderFeature
11from music_assistant_models.errors import MusicAssistantError
12
13from music_assistant.providers.sonic_similarity import SonicSimilarityPlugin, setup
14from music_assistant.providers.sonic_similarity import clap_index as clap_index_module
15from music_assistant.providers.sonic_similarity.similarity import ScoredCandidate
16from tests.providers.sonic_similarity.conftest import make_track
17
18if TYPE_CHECKING:
19 from collections.abc import Callable
20 from typing import Any
21
22
23def _make_config(*, enable_text_search: bool) -> MagicMock:
24 """Build a ProviderConfig mock with the keys setup()/__init__ read."""
25 config = MagicMock()
26 config_values = {
27 "log_level": "GLOBAL",
28 "aa_provider_domain": "sonic_analysis",
29 "enable_clap_index": False,
30 "enable_text_search": enable_text_search,
31 "enable_discover_row": True,
32 "discover_preset": "discover",
33 "discover_diversity": 0.2,
34 }
35 config.get_value = lambda key: config_values.get(key)
36 return config
37
38
39def _make_manifest() -> MagicMock:
40 """Build a ProviderManifest mock with the attributes the base Provider reads."""
41 manifest = MagicMock()
42 manifest.instance_id = "test-instance-id"
43 manifest.domain = "sonic_similarity"
44 return manifest
45
46
47class TestSetupConditionalSearch:
48 """setup() advertises ProviderFeature.SEARCH only when text search is enabled."""
49
50 @pytest.mark.asyncio
51 async def test_search_feature_present_when_text_search_enabled(
52 self, mock_mass: MagicMock
53 ) -> None:
54 """CONF_ENABLE_TEXT_SEARCH=True â plugin advertises SEARCH."""
55 # setup() reads the stored enable_text_search value directly from mass.config
56 mock_mass.config.get_raw_provider_config_value = MagicMock(return_value=True)
57 plugin = await setup(mock_mass, _make_manifest(), _make_config(enable_text_search=True))
58 assert ProviderFeature.SEARCH in plugin.supported_features
59
60 @pytest.mark.asyncio
61 async def test_search_feature_absent_when_text_search_disabled(
62 self, mock_mass: MagicMock
63 ) -> None:
64 """CONF_ENABLE_TEXT_SEARCH=False â plugin does not advertise SEARCH."""
65 # setup() reads the stored enable_text_search value directly from mass.config
66 mock_mass.config.get_raw_provider_config_value = MagicMock(return_value=False)
67 plugin = await setup(mock_mass, _make_manifest(), _make_config(enable_text_search=False))
68 assert ProviderFeature.SEARCH not in plugin.supported_features
69
70
71class TestTextEncoderWarmsLazily:
72 """
73 The GPT2 text encoder is not warmed at load; the first search() warms it lazily.
74
75 Warming the ~500MB encoder eagerly at startup defeats the point of an opt-in
76 feature, so loaded_in_mass leaves it cold and the first cold search() kicks off
77 a one-time background warm (see TestSearch.test_schedules_warm_when_encoder_cold).
78 """
79
80 @pytest.mark.asyncio
81 async def test_no_warm_task_at_load_when_text_search_enabled(
82 self, mock_mass: MagicMock, monkeypatch: pytest.MonkeyPatch
83 ) -> None:
84 """text_search_enabled â loaded_in_mass does NOT warm the encoder (it loads on first query)."""
85
86 async def _noop_load(_self: Any) -> None:
87 return None
88
89 # text_search auto-enables CLAP; stub its load() so we don't touch disk.
90 monkeypatch.setattr(clap_index_module.ClapIndex, "load", _noop_load)
91
92 plugin = await setup(mock_mass, _make_manifest(), _make_config(enable_text_search=True))
93 assert isinstance(plugin, SonicSimilarityPlugin)
94 await plugin.loaded_in_mass()
95
96 mock_mass.create_task.assert_not_called()
97
98 @pytest.mark.asyncio
99 async def test_no_warm_task_when_text_search_disabled(self, mock_mass: MagicMock) -> None:
100 """text_search_enabled=False â no background warm scheduled."""
101 plugin = await setup(mock_mass, _make_manifest(), _make_config(enable_text_search=False))
102 await plugin.loaded_in_mass()
103
104 mock_mass.create_task.assert_not_called()
105
106
107def _make_tensor_chain(vector: np.ndarray) -> MagicMock:
108 """Build a CLAP-tensor-like MagicMock whose .detach().cpu().numpy() yields vector."""
109 chain = MagicMock()
110 chain.detach.return_value = chain
111 chain.cpu.return_value = chain
112 chain.numpy.return_value = vector
113 return chain
114
115
116def _make_mock_encoder(vector: np.ndarray) -> MagicMock:
117 """Build a CLAP-like encoder whose get_text_embeddings([q]) yields [tensor_chain]."""
118 encoder = MagicMock()
119 encoder.get_text_embeddings.return_value = [_make_tensor_chain(vector)]
120 return encoder
121
122
123class TestSearch:
124 """Tests for SonicSimilarityPlugin.search (the ProviderFeature.SEARCH dispatcher)."""
125
126 @pytest.mark.asyncio
127 async def test_returns_empty_when_track_not_in_media_types(
128 self, make_plugin: Callable[..., Any]
129 ) -> None:
130 """We only search tracks; non-track media_types return empty SearchResults."""
131 plugin = make_plugin(clap_enabled=True)
132 plugin._clap_index.__len__ = MagicMock(return_value=5)
133
134 result = await plugin.search("disco", [MediaType.ALBUM, MediaType.ARTIST])
135
136 assert list(result.tracks) == []
137 # The encoder must not even be touched when no track type was requested.
138 plugin._clap_index.search.assert_not_called()
139
140 @pytest.mark.asyncio
141 async def test_returns_empty_when_no_clap_index(self, make_plugin: Callable[..., Any]) -> None:
142 """Without a CLAP index attached the search short-circuits."""
143 plugin = make_plugin() # clap_enabled=False â _clap_index stays None
144
145 result = await plugin.search("disco", [MediaType.TRACK])
146
147 assert list(result.tracks) == []
148
149 @pytest.mark.asyncio
150 async def test_returns_empty_when_clap_index_is_empty(
151 self, make_plugin: Callable[..., Any]
152 ) -> None:
153 """An empty CLAP index (len == 0) short-circuits without encoding."""
154 plugin = make_plugin(clap_enabled=True) # default len == 0
155
156 result = await plugin.search("disco", [MediaType.TRACK])
157
158 assert list(result.tracks) == []
159 plugin._clap_index.search.assert_not_called()
160
161 @pytest.mark.asyncio
162 async def test_schedules_warm_when_encoder_cold(self, make_plugin: Callable[..., Any]) -> None:
163 """
164 A cold encoder short-circuits to empty but kicks off a one-time background warm.
165
166 The encoder loads on first query rather than at startup, but the load runs off
167 the request path (via mass.create_task) so the global SEARCH dispatcher never
168 blocks on the ~500MB GPT2 download.
169 """
170 plugin = make_plugin(clap_enabled=True)
171 plugin._clap_index.__len__ = MagicMock(return_value=5)
172 # Sentinel: search() must not load the encoder synchronously on the request path.
173 plugin._load_text_encoder = MagicMock(
174 side_effect=RuntimeError("must not load synchronously")
175 )
176
177 result = await plugin.search("disco", [MediaType.TRACK])
178
179 assert list(result.tracks) == []
180 plugin._load_text_encoder.assert_not_called()
181 plugin._clap_index.search.assert_not_called()
182 plugin.mass.create_task.assert_called_once_with(
183 plugin._get_text_encoder, task_id="sonic_similarity_text_encoder_warm"
184 )
185
186 @pytest.mark.asyncio
187 async def test_happy_path_returns_resolved_tracks(
188 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
189 ) -> None:
190 """Encoded query â CLAP matches â resolved Track objects in SearchResults.tracks."""
191 plugin = make_plugin(clap_enabled=True)
192 plugin._clap_index.__len__ = MagicMock(return_value=5)
193 vector = np.full((1024,), 0.1, dtype=np.float32)
194 plugin._text_encoder = _make_mock_encoder(vector)
195 plugin._clap_index.search = AsyncMock(
196 return_value=[
197 ScoredCandidate("track_a", "spotify", 0.1),
198 ScoredCandidate("track_b", "tidal", 0.2),
199 ]
200 )
201 track_a = make_track("track_a", provider="spotify", name="A")
202 track_b = make_track("track_b", provider="tidal", name="B")
203 mock_mass.music.tracks.get.side_effect = [track_a, track_b]
204
205 result = await plugin.search("disco", [MediaType.TRACK], limit=5)
206
207 assert list(result.tracks) == [track_a, track_b]
208 plugin._clap_index.search.assert_awaited_once()
209
210 @pytest.mark.asyncio
211 async def test_drops_resolves_that_raise(
212 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
213 ) -> None:
214 """An unresolvable item is silently dropped; the rest pass through."""
215 plugin = make_plugin(clap_enabled=True)
216 plugin._clap_index.__len__ = MagicMock(return_value=5)
217 vector = np.full((1024,), 0.1, dtype=np.float32)
218 plugin._text_encoder = _make_mock_encoder(vector)
219 plugin._clap_index.search = AsyncMock(
220 return_value=[
221 ScoredCandidate("track_a", "spotify", 0.1),
222 ScoredCandidate("track_b", "tidal", 0.2),
223 ]
224 )
225 track_b = make_track("track_b", provider="tidal", name="B")
226
227 async def _fake_get(item_id: str, _provider: str) -> MagicMock:
228 if item_id == "track_a":
229 raise MusicAssistantError("nope")
230 return track_b
231
232 mock_mass.music.tracks.get.side_effect = _fake_get
233
234 result = await plugin.search("disco", [MediaType.TRACK], limit=5)
235
236 assert list(result.tracks) == [track_b]
237
238 @pytest.mark.asyncio
239 async def test_passes_limit_to_clap_index_search(
240 self, make_plugin: Callable[..., Any], mock_mass: MagicMock
241 ) -> None:
242 """The limit kwarg is forwarded as the k argument to CLAP index search."""
243 plugin = make_plugin(clap_enabled=True)
244 plugin._clap_index.__len__ = MagicMock(return_value=20)
245 vector = np.full((1024,), 0.1, dtype=np.float32)
246 plugin._text_encoder = _make_mock_encoder(vector)
247 plugin._clap_index.search = AsyncMock(return_value=[])
248 mock_mass.music.tracks.get = AsyncMock()
249
250 await plugin.search("disco", [MediaType.TRACK], limit=7)
251
252 # The CLAP index search is called positionally as (embedding, k).
253 call = plugin._clap_index.search.await_args
254 assert call.args[1] == 7
255