/
/
/
1"""Tests for the post load steps that run once a provider is registered."""
2
3from __future__ import annotations
4
5import asyncio
6from typing import cast
7from unittest.mock import AsyncMock, MagicMock
8
9from music_assistant_models.enums import ProviderType
10
11from music_assistant.mass import MusicAssistant
12from music_assistant.models.plugin import PluginProvider
13from tests.common import use_real_create_task
14
15
16class _Provider(PluginProvider):
17 """Provider whose post load step raises what its test asks for."""
18
19 post_load_error: BaseException | None = None
20
21 async def loaded_in_mass(self) -> None:
22 """Fail the way a provider does when a post load step hits a problem."""
23 if self.post_load_error:
24 raise self.post_load_error
25
26
27def _mass() -> MusicAssistant:
28 """Return a bare MusicAssistant (bypassing __init__) able to register a provider."""
29 mass = object.__new__(MusicAssistant)
30 mass._providers = {}
31 mass._provider_ready_events = {}
32 mass.cache = MagicMock()
33 mass.config = MagicMock()
34 mass.discovery = MagicMock()
35 mass.signal_event = MagicMock() # type: ignore[method-assign]
36 mass._update_available_providers_cache = AsyncMock() # type: ignore[method-assign]
37 mass.run_provider_discovery = AsyncMock() # type: ignore[method-assign]
38 use_real_create_task(mass)
39 return mass
40
41
42def _provider(mass: MusicAssistant, post_load_error: BaseException | None = None) -> _Provider:
43 """Return a provider instance for the 'test' domain."""
44 manifest = MagicMock()
45 manifest.domain = "test"
46 manifest.name = "Test Provider"
47 manifest.type = ProviderType.PLUGIN
48 config = MagicMock()
49 config.instance_id = "test--1"
50 config.name = "Test"
51 config.get_value.return_value = "GLOBAL"
52 provider = _Provider(mass, manifest, config)
53 provider.post_load_error = post_load_error
54 return provider
55
56
57async def test_post_load_runs_every_step() -> None:
58 """A provider that loaded runs all of its post load steps."""
59 mass = _mass()
60 provider = _provider(mass)
61
62 await mass._register_loaded_provider(provider, provider.config)
63 # the post load steps run as a task of their own, so let it run to completion
64 await asyncio.sleep(0)
65
66 assert provider.initialized.is_set()
67 assert mass.get_provider_ready_event("test").is_set()
68 cast("AsyncMock", mass.run_provider_discovery).assert_awaited_once_with("test--1")
69
70
71async def test_failing_post_load_still_runs_the_remaining_steps() -> None:
72 """A failed post load step must not leave the waiters of a provider hanging."""
73 mass = _mass()
74 provider = _provider(mass, RuntimeError("post load failed"))
75
76 await mass._register_loaded_provider(provider, provider.config)
77 await asyncio.sleep(0)
78
79 assert provider.available is True
80 assert provider.initialized.is_set()
81 assert mass.get_provider_ready_event("test").is_set()
82 # discovery and the default name are unrelated to whatever failed, so they still run
83 cast("AsyncMock", mass.run_provider_discovery).assert_awaited_once_with("test--1")
84 cast("MagicMock", mass.config).set_provider_default_name.assert_called_once_with(
85 "test--1", "Test Provider"
86 )
87
88
89async def test_cancelled_post_load_reports_nothing() -> None:
90 """A provider that is torn down mid load must not be announced as ready."""
91 mass = _mass()
92 provider = _provider(mass, asyncio.CancelledError())
93
94 await mass._register_loaded_provider(provider, provider.config)
95 await asyncio.sleep(0)
96
97 assert not provider.initialized.is_set()
98 assert not mass.get_provider_ready_event("test").is_set()
99 cast("AsyncMock", mass.run_provider_discovery).assert_not_awaited()
100