/
/
/
1"""Tests for reloading providers when the streamserver network changes."""
2
3from __future__ import annotations
4
5import asyncio
6from typing import cast
7from unittest.mock import AsyncMock, MagicMock
8
9import pytest
10
11from music_assistant.controllers.streams.controller import StreamsController
12from music_assistant.models.provider import Provider
13
14BIND_IP = "0.0.0.0"
15PUBLISH_IP = "192.168.1.5"
16PUBLISH_PORT = 8097
17PUBLISH_ADDRESSES = ["192.168.1.5", "10.0.0.5"]
18
19
20def _provider(instance_id: str, *, follows_network: bool) -> MagicMock:
21 """Build a stand-in provider that does or does not capture the streamserver network."""
22 provider = MagicMock()
23 provider.instance_id = instance_id
24 provider.reload_on_streams_network_change = follows_network
25 return provider
26
27
28def _provider_config(instance_id: str) -> MagicMock:
29 """Build the stored config a provider would be reloaded from."""
30 config = MagicMock()
31 config.name = instance_id
32 config.domain = instance_id
33 return config
34
35
36def _streams_controller(*providers: MagicMock) -> StreamsController:
37 """Build a streams controller with only the network state populated."""
38 mass = MagicMock()
39 mass.config.get_raw_core_config_value.return_value = "GLOBAL"
40 mass.providers = list(providers)
41 mass.config.get_provider_config = AsyncMock(side_effect=_provider_config)
42 mass.load_provider_config = AsyncMock()
43 controller = StreamsController(mass)
44 controller._bind_ip = BIND_IP
45 controller.publish_ip = PUBLISH_IP
46 controller.publish_port = PUBLISH_PORT
47 controller._publish_addresses = list(PUBLISH_ADDRESSES)
48 return controller
49
50
51def _mass(controller: StreamsController) -> MagicMock:
52 """Return the mocked server object the controller was built with."""
53 return cast("MagicMock", controller.mass)
54
55
56def _reloaded(controller: StreamsController) -> list[str]:
57 """Return the names of the providers that were reloaded."""
58 return [call.args[0].name for call in _mass(controller).load_provider_config.await_args_list]
59
60
61@pytest.mark.asyncio
62async def test_startup_network_is_only_recorded() -> None:
63 """The network at startup is the baseline, not a change to react to."""
64 controller = _streams_controller(_provider("sendspin", follows_network=True))
65 await controller._reload_network_dependent_providers()
66 assert _reloaded(controller) == []
67
68
69@pytest.mark.asyncio
70async def test_reload_without_network_change_reloads_nothing() -> None:
71 """Any other streams setting is reloadable on its own, so nobody else is disturbed."""
72 controller = _streams_controller(_provider("sendspin", follows_network=True))
73 await controller._reload_network_dependent_providers()
74 await controller._reload_network_dependent_providers()
75 assert _reloaded(controller) == []
76
77
78@pytest.mark.asyncio
79@pytest.mark.parametrize(
80 ("attribute", "new_value"),
81 [
82 ("_bind_ip", "192.168.1.9"),
83 ("publish_ip", "10.45.0.20"),
84 ("publish_port", 8098),
85 # narrowing "auto" to one literal address leaves publish_ip itself untouched
86 ("_publish_addresses", ["192.168.1.5"]),
87 ],
88)
89async def test_changed_network_reloads_only_dependent_providers(
90 attribute: str, new_value: str | int | list[str]
91) -> None:
92 """Every part of the network is captured while a provider loads, so each must count."""
93 controller = _streams_controller(
94 _provider("sendspin", follows_network=True),
95 _provider("spotify", follows_network=False),
96 )
97 await controller._reload_network_dependent_providers()
98 setattr(controller, attribute, new_value)
99 await controller._reload_network_dependent_providers()
100 assert _reloaded(controller) == ["sendspin"]
101
102
103@pytest.mark.asyncio
104async def test_applied_network_is_not_reloaded_again() -> None:
105 """A later reload compares against the applied network, not the one seen at startup."""
106 controller = _streams_controller(_provider("sendspin", follows_network=True))
107 await controller._reload_network_dependent_providers()
108 controller.publish_ip = "10.45.0.20"
109 await controller._reload_network_dependent_providers()
110 await controller._reload_network_dependent_providers()
111 assert _reloaded(controller) == ["sendspin"]
112
113
114@pytest.mark.asyncio
115async def test_interrupted_run_does_not_mark_the_network_applied() -> None:
116 """A second config change cancels the run, so the next reload goes over them again."""
117 controller = _streams_controller(_provider("sendspin", follows_network=True))
118 await controller._reload_network_dependent_providers()
119 controller.publish_ip = "10.45.0.20"
120 _mass(controller).load_provider_config = AsyncMock(side_effect=[asyncio.CancelledError, None])
121 with pytest.raises(asyncio.CancelledError):
122 await controller._reload_network_dependent_providers()
123 await controller._reload_network_dependent_providers()
124 # the cancelled attempt plus the one the still-unapplied network triggered
125 assert _reloaded(controller) == ["sendspin", "sendspin"]
126
127
128def test_providers_do_not_follow_the_network_by_default() -> None:
129 """Only a provider that opts in is ever disturbed by a network change."""
130 assert Provider.reload_on_streams_network_change is False
131
132
133@pytest.mark.asyncio
134async def test_unreadable_config_does_not_block_the_others() -> None:
135 """A provider removed while the loop runs no longer has a config to reload it from."""
136 controller = _streams_controller(
137 _provider("snapcast", follows_network=True),
138 _provider("sendspin", follows_network=True),
139 )
140 _mass(controller).config.get_provider_config = AsyncMock(
141 side_effect=[KeyError("gone"), _provider_config("sendspin")]
142 )
143 await controller._reload_network_dependent_providers()
144 controller.publish_ip = "10.45.0.20"
145 await controller._reload_network_dependent_providers()
146 assert _reloaded(controller) == ["sendspin"]
147
148
149@pytest.mark.asyncio
150async def test_failed_reload_still_marks_the_network_applied() -> None:
151 """A provider that stays broken must not re-bounce the others on every later reload."""
152 controller = _streams_controller(_provider("sendspin", follows_network=True))
153 _mass(controller).load_provider_config = AsyncMock(side_effect=RuntimeError("boom"))
154 await controller._reload_network_dependent_providers()
155 controller.publish_ip = "10.45.0.20"
156 await controller._reload_network_dependent_providers()
157 await controller._reload_network_dependent_providers()
158 assert _reloaded(controller) == ["sendspin"]
159
160
161@pytest.mark.asyncio
162async def test_failing_reload_does_not_block_the_others() -> None:
163 """One provider that cannot come back up must not strand the rest on the old network."""
164 controller = _streams_controller(
165 _provider("snapcast", follows_network=True),
166 _provider("sendspin", follows_network=True),
167 )
168 _mass(controller).load_provider_config = AsyncMock(side_effect=[RuntimeError("boom"), None])
169 await controller._reload_network_dependent_providers()
170 controller.publish_ip = "10.45.0.20"
171 await controller._reload_network_dependent_providers()
172 assert _reloaded(controller) == ["snapcast", "sendspin"]
173