/
/
/
1"""Tests for the global search on the music controller."""
2
3from __future__ import annotations
4
5import asyncio
6import logging
7import time
8from typing import TYPE_CHECKING, Any, cast
9from unittest.mock import AsyncMock, Mock
10
11from music_assistant_models.enums import MediaType, ProviderFeature
12from music_assistant_models.errors import MusicAssistantError
13from music_assistant_models.media_items import (
14 Artist,
15 Genre,
16 ProviderMapping,
17 SearchResults,
18 Track,
19)
20
21from music_assistant.controllers.music import MusicController
22from music_assistant.controllers.music.constants import (
23 SEARCH_CACHE_EXPIRATION_LOCAL_PROVIDER,
24 SEARCH_CACHE_EXPIRATION_STREAMING_PROVIDER,
25)
26from music_assistant.models.music_provider import MusicProvider
27
28if TYPE_CHECKING:
29 from collections.abc import Callable
30
31 import pytest
32
33
34def _make_track(item_id: str, provider: str, name: str) -> Track:
35 """Return a minimal Track for the given provider."""
36 return Track(
37 item_id=item_id,
38 provider=provider,
39 name=name,
40 provider_mappings={
41 ProviderMapping(
42 item_id=item_id,
43 provider_domain=provider,
44 provider_instance=provider,
45 )
46 },
47 )
48
49
50def _make_library_artist(name: str, mapped_provider: str) -> Artist:
51 """Return a library Artist with a provider mapping to the given provider."""
52 return Artist(
53 item_id="lib1",
54 provider="library",
55 name=name,
56 provider_mappings={
57 ProviderMapping(
58 item_id="artist1",
59 provider_domain=mapped_provider,
60 provider_instance=mapped_provider,
61 )
62 },
63 )
64
65
66def _make_search_provider(instance_id: str, domain: str | None = None) -> Mock:
67 """Return a mocked music provider that supports search."""
68 prov = Mock(spec=MusicProvider)
69 prov.instance_id = instance_id
70 prov.domain = domain or instance_id
71 prov.name = instance_id
72 prov.supported_features = {ProviderFeature.SEARCH}
73 prov.is_streaming_provider = True
74 prov.search = AsyncMock(return_value=SearchResults())
75 return prov
76
77
78def _make_controller(providers: list[Mock]) -> MusicController:
79 """Return a music controller wired to the given mocked providers."""
80 mass = Mock()
81 mass.cache.get = AsyncMock(return_value=None)
82 mass.cache.set = AsyncMock(return_value=None)
83 mass.get_providers_supporting_feature.return_value = []
84 mass.get_provider = Mock(
85 side_effect=lambda instance_id, **_kwargs: next(
86 (p for p in providers if p.instance_id == instance_id), None
87 )
88 )
89 # mimic the task_id deduplication behavior of the real create_task implementation
90 tracked_tasks: dict[str, asyncio.Task[Any]] = {}
91
92 def _create_task(
93 target: Any, *_args: Any, task_id: str | None = None, **_kwargs: Any
94 ) -> asyncio.Task[Any]:
95 if task_id and (existing := tracked_tasks.get(task_id)) and not existing.done():
96 target.close()
97 return existing
98 task = asyncio.get_running_loop().create_task(target)
99
100 def _task_done(_task: asyncio.Task[Any]) -> None:
101 if task_id:
102 tracked_tasks.pop(task_id, None)
103 if not _task.cancelled():
104 _task.exception()
105
106 task.add_done_callback(_task_done)
107 if task_id:
108 tracked_tasks[task_id] = task
109 return task
110
111 mass.create_task = Mock(side_effect=_create_task)
112 controller = MusicController.__new__(MusicController)
113 controller.mass = mass
114 controller.domain = "music"
115 controller.logger = logging.getLogger(__name__)
116 controller.get_unique_providers = Mock( # type: ignore[method-assign]
117 return_value=[p.instance_id for p in providers]
118 )
119 controller.search_library = AsyncMock(return_value=SearchResults()) # type: ignore[method-assign]
120 return controller
121
122
123def _cache_writes(controller: MusicController, provider: str) -> list[Any]:
124 """Return the cache.set calls made for the given provider."""
125 cache_set = cast("AsyncMock", controller.mass.cache.set)
126 return [call for call in cache_set.await_args_list if call.kwargs.get("provider") == provider]
127
128
129async def _wait_for(condition: Callable[[], bool], timeout: float = 1.0) -> None:
130 """Wait until the given condition is true."""
131 async with asyncio.timeout(timeout):
132 while not condition():
133 await asyncio.sleep(0.01)
134
135
136async def test_search_provider_returns_none_on_provider_error() -> None:
137 """A provider error during search yields None instead of raising."""
138 prov = _make_search_provider("prov_a")
139 controller = _make_controller([prov])
140 for error in (MusicAssistantError("rate limited"), ValueError("unexpected")):
141 prov.search.side_effect = error
142 result = await controller._search_provider("query", "prov_a", [MediaType.TRACK])
143 assert result is None
144
145
146async def test_global_search_returns_partial_results_when_provider_fails() -> None:
147 """One failing provider must not break the entire global search."""
148 prov_ok = _make_search_provider("prov_ok")
149 prov_ok.search.return_value = SearchResults(
150 tracks=[_make_track("track1", "prov_ok", "My Song")]
151 )
152 prov_bad = _make_search_provider("prov_bad")
153 prov_bad.search.side_effect = MusicAssistantError("provider down")
154 controller = _make_controller([prov_ok, prov_bad])
155
156 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
157
158 assert [track.item_id for track in result.tracks] == ["track1"]
159 prov_bad.search.assert_awaited_once()
160 # an incomplete result may not be cached so the failed provider is retried
161 assert not _cache_writes(controller, "music")
162
163
164async def test_global_search_soft_timeout_returns_partial_and_caches_late(
165 monkeypatch: pytest.MonkeyPatch,
166) -> None:
167 """A slow provider is not awaited beyond the soft timeout but still caches its result."""
168 monkeypatch.setattr(
169 "music_assistant.controllers.music.controller.SEARCH_PROVIDER_SOFT_TIMEOUT", 0.1
170 )
171 prov_fast = _make_search_provider("prov_fast")
172 prov_fast.search.return_value = SearchResults(
173 tracks=[_make_track("track1", "prov_fast", "My Song")]
174 )
175 prov_slow = _make_search_provider("prov_slow")
176 slow_results = SearchResults(tracks=[_make_track("track2", "prov_slow", "My Song")])
177
178 async def _slow_search(*_args: Any, **_kwargs: Any) -> SearchResults:
179 await asyncio.sleep(0.5)
180 return slow_results
181
182 prov_slow.search.side_effect = _slow_search
183 controller = _make_controller([prov_fast, prov_slow])
184
185 start = time.monotonic()
186 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
187 duration = time.monotonic() - start
188
189 # the search returned within the soft timeout with the results of the fast provider
190 assert duration < 0.4
191 assert [track.item_id for track in result.tracks] == ["track1"]
192 assert not _cache_writes(controller, "music")
193 # the slow provider search completes in the background and caches its result
194 await _wait_for(lambda: bool(_cache_writes(controller, "prov_slow")))
195 late_writes = _cache_writes(controller, "prov_slow")
196 assert len(late_writes) == 1
197 assert late_writes[0].kwargs["data"] == slow_results.to_dict()
198
199
200async def test_global_search_hard_timeout_aborts_provider_search(
201 monkeypatch: pytest.MonkeyPatch,
202) -> None:
203 """A provider search that exceeds the hard timeout is aborted and never cached."""
204 monkeypatch.setattr(
205 "music_assistant.controllers.music.controller.SEARCH_PROVIDER_SOFT_TIMEOUT", 0.05
206 )
207 monkeypatch.setattr(
208 "music_assistant.controllers.music.controller.SEARCH_PROVIDER_HARD_TIMEOUT", 0.1
209 )
210 prov_hung = _make_search_provider("prov_hung")
211 search_aborted = asyncio.Event()
212
213 async def _hung_search(*_args: Any, **_kwargs: Any) -> SearchResults:
214 try:
215 await asyncio.sleep(5)
216 finally:
217 search_aborted.set()
218 return SearchResults()
219
220 prov_hung.search.side_effect = _hung_search
221 controller = _make_controller([prov_hung])
222
223 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
224
225 assert result.tracks == []
226 # the hard timeout cancels the provider search and nothing is cached
227 await _wait_for(search_aborted.is_set)
228 cache_set = cast("AsyncMock", controller.mass.cache.set)
229 assert not cache_set.await_args_list
230
231
232async def test_search_provider_cache_hit_avoids_provider_call() -> None:
233 """A cached provider result is served without hitting the provider again."""
234 prov = _make_search_provider("prov_a")
235 controller = _make_controller([prov])
236 cached_results = SearchResults(tracks=[_make_track("track1", "prov_a", "My Song")])
237
238 async def _fake_cache_get(*, provider: str = "default", **_kwargs: Any) -> Any:
239 return cached_results if provider == "prov_a" else None
240
241 controller.mass.cache.get = AsyncMock(side_effect=_fake_cache_get) # type: ignore[method-assign]
242
243 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
244
245 assert [track.item_id for track in result.tracks] == ["track1"]
246 prov.search.assert_not_awaited()
247
248
249async def test_empty_search_query_skips_all_providers() -> None:
250 """An empty or whitespace-only query returns empty results without querying providers."""
251 prov = _make_search_provider("prov_a")
252 controller = _make_controller([prov])
253
254 for search_query in ("", " ", "\n\t"):
255 result = await controller.search(search_query, media_types=[MediaType.TRACK], limit=5)
256
257 assert result == SearchResults()
258 prov.search.assert_not_awaited()
259 cast("AsyncMock", controller.search_library).assert_not_awaited()
260 cast("AsyncMock", controller.mass.cache.get).assert_not_awaited()
261
262
263async def test_search_provider_cache_expiration_by_provider_type() -> None:
264 """Streaming provider results are cached longer than local provider results."""
265 prov_stream = _make_search_provider("prov_stream")
266 prov_local = _make_search_provider("prov_local")
267 prov_local.is_streaming_provider = False
268 controller = _make_controller([prov_stream, prov_local])
269
270 await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
271
272 stream_writes = _cache_writes(controller, "prov_stream")
273 local_writes = _cache_writes(controller, "prov_local")
274 assert stream_writes[0].kwargs["expiration"] == SEARCH_CACHE_EXPIRATION_STREAMING_PROVIDER
275 assert local_writes[0].kwargs["expiration"] == SEARCH_CACHE_EXPIRATION_LOCAL_PROVIDER
276
277
278async def test_search_provider_without_streaming_attribute_gets_local_expiration() -> None:
279 """Providers lacking is_streaming_provider (e.g. plugins) get the short cache expiration."""
280 prov = _make_search_provider("prov_plugin")
281 del prov.is_streaming_provider
282 prov.search.return_value = SearchResults(
283 tracks=[_make_track("track1", "prov_plugin", "My Song")]
284 )
285 controller = _make_controller([prov])
286
287 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
288
289 assert [track.item_id for track in result.tracks] == ["track1"]
290 writes = _cache_writes(controller, "prov_plugin")
291 assert writes[0].kwargs["expiration"] == SEARCH_CACHE_EXPIRATION_LOCAL_PROVIDER
292
293
294async def test_search_cache_write_failure_does_not_break_search() -> None:
295 """A failing cache write still returns the provider results."""
296 prov = _make_search_provider("prov_a")
297 prov.search.return_value = SearchResults(tracks=[_make_track("track1", "prov_a", "My Song")])
298 controller = _make_controller([prov])
299 controller.mass.cache.set = AsyncMock(side_effect=RuntimeError("cache boom")) # type: ignore[method-assign]
300
301 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
302
303 assert [track.item_id for track in result.tracks] == ["track1"]
304
305
306async def test_concurrent_identical_searches_share_one_provider_call() -> None:
307 """Identical concurrent searches are coalesced into a single provider call."""
308 prov = _make_search_provider("prov_a")
309
310 async def _slowish_search(*_args: Any, **_kwargs: Any) -> SearchResults:
311 await asyncio.sleep(0.05)
312 return SearchResults(tracks=[_make_track("track1", "prov_a", "My Song")])
313
314 prov.search.side_effect = _slowish_search
315 controller = _make_controller([prov])
316
317 result1, result2 = await asyncio.gather(
318 controller.search("My Song", media_types=[MediaType.TRACK], limit=5),
319 controller.search("My Song", media_types=[MediaType.TRACK], limit=5),
320 )
321
322 assert [track.item_id for track in result1.tracks] == ["track1"]
323 assert [track.item_id for track in result2.tracks] == ["track1"]
324 prov.search.assert_awaited_once()
325
326
327async def test_search_dedup_filters_provider_items_already_in_library() -> None:
328 """Provider items that map to a library item are filtered from the results."""
329 library_track = Track(
330 item_id="lib1",
331 provider="library",
332 name="Library Song",
333 provider_mappings={
334 ProviderMapping(
335 item_id="track1",
336 provider_domain="prov_a",
337 provider_instance="prov_a",
338 )
339 },
340 )
341 prov = _make_search_provider("prov_a")
342 prov.search.return_value = SearchResults(
343 tracks=[
344 _make_track("track1", "prov_a", "My Song"),
345 _make_track("track2", "prov_a", "My Song 2"),
346 ]
347 )
348 controller = _make_controller([prov])
349 controller.search_library = AsyncMock( # type: ignore[method-assign]
350 return_value=SearchResults(tracks=[library_track])
351 )
352
353 result = await controller.search("My Song", media_types=[MediaType.TRACK], limit=5)
354
355 assert [track.item_id for track in result.tracks] == ["lib1", "track2"]
356
357
358async def test_search_includes_genre_results_from_library() -> None:
359 """Genre results from the library search end up in the combined search result."""
360 library_genre = Genre(
361 item_id="genre1", provider="library", name="Rock", provider_mappings=set()
362 )
363 prov = _make_search_provider("prov_a")
364 controller = _make_controller([prov])
365 controller.search_library = AsyncMock( # type: ignore[method-assign]
366 return_value=SearchResults(genres=[library_genre])
367 )
368
369 result = await controller.search("Rock", media_types=[MediaType.GENRE], limit=5)
370
371 assert [genre.item_id for genre in result.genres] == ["genre1"]
372
373
374async def test_search_providers_param_restricts_search() -> None:
375 """The providers param restricts the search to the given providers."""
376 prov_a = _make_search_provider("prov_a")
377 prov_b = _make_search_provider("spotify--abc", domain="spotify")
378 prov_b.search.return_value = SearchResults(
379 tracks=[_make_track("track1", "spotify--abc", "My Song")]
380 )
381 controller = _make_controller([prov_a, prov_b])
382 controller.search_library = AsyncMock( # type: ignore[method-assign]
383 return_value=SearchResults(tracks=[_make_track("lib1", "library", "My Song")])
384 )
385
386 # domain match: only the matching provider is searched, library results excluded
387 result = await controller.search(
388 "My Song", media_types=[MediaType.TRACK], limit=5, providers=["spotify"]
389 )
390 assert [track.item_id for track in result.tracks] == ["track1"]
391 prov_a.search.assert_not_awaited()
392 prov_b.search.assert_awaited_once()
393
394 # the special value "library" selects the library only
395 prov_b.search.reset_mock()
396 result = await controller.search(
397 "My Song", media_types=[MediaType.TRACK], limit=5, providers=["library"]
398 )
399 assert [track.item_id for track in result.tracks] == ["lib1"]
400 prov_a.search.assert_not_awaited()
401 prov_b.search.assert_not_awaited()
402
403
404async def test_search_library_only_maps_to_library_provider() -> None:
405 """The deprecated library_only flag behaves as providers=["library"]."""
406 prov = _make_search_provider("prov_a")
407 controller = _make_controller([prov])
408 controller.search_library = AsyncMock( # type: ignore[method-assign]
409 return_value=SearchResults(tracks=[_make_track("lib1", "library", "My Song")])
410 )
411
412 result = await controller.search(
413 "My Song", media_types=[MediaType.TRACK], limit=5, library_only=True
414 )
415
416 assert [track.item_id for track in result.tracks] == ["lib1"]
417 prov.search.assert_not_awaited()
418
419
420async def test_search_exact_library_match_skips_provider_media_types() -> None:
421 """A (near) exact library match skips searching mapped providers for that media type."""
422 prov = _make_search_provider("prov_a")
423 controller = _make_controller([prov])
424 controller.search_library = AsyncMock( # type: ignore[method-assign]
425 return_value=SearchResults(artists=[_make_library_artist("Nirvana", "prov_a")])
426 )
427
428 # the only requested media type is covered by the library: skip the provider
429 result = await controller.search("Nirvana", media_types=[MediaType.ARTIST], limit=5)
430 assert [artist.item_id for artist in result.artists] == ["lib1"]
431 prov.search.assert_not_awaited()
432
433 # other media types are still searched on the provider
434 await controller.search("Nirvana", media_types=[MediaType.ARTIST, MediaType.TRACK], limit=5)
435 prov.search.assert_awaited_once()
436 assert prov.search.await_args.args[1] == [MediaType.TRACK]
437
438
439async def test_search_exact_match_shortcut_skipped_for_explicit_providers() -> None:
440 """An explicit providers selection always searches those providers."""
441 prov = _make_search_provider("prov_a")
442 controller = _make_controller([prov])
443 controller.search_library = AsyncMock( # type: ignore[method-assign]
444 return_value=SearchResults(artists=[_make_library_artist("Nirvana", "prov_a")])
445 )
446
447 await controller.search(
448 "Nirvana", media_types=[MediaType.ARTIST], limit=5, providers=["prov_a"]
449 )
450
451 prov.search.assert_awaited_once()
452