/
/
1"""Tests for music_assistant.providers.snapcast.ma_stream._register_tcp_server_source."""
2
3from __future__ import annotations
4
5import asyncio
6from typing import TYPE_CHECKING
7from unittest.mock import MagicMock
8
9import pytest
10from music_assistant_models.enums import ContentType
11
12from music_assistant.providers.snapcast.constants import (
13 DEFAULT_SNAPCAST_FORMAT,
14 snapcast_sampleformat_query,
15 snapcast_stream_format,
16)
17from music_assistant.providers.snapcast.ma_stream import SnapcastMAStream
18
19if TYPE_CHECKING:
20 from .conftest import FakeSnapserver
21
22
23def _make_stream(
24 provider: MagicMock, name: str = "Music Assistant - testhash (announcement)"
25) -> SnapcastMAStream:
26 """
27 Build a SnapcastMAStream instance directly.
28
29 Bypasses the constructor's media handling (we only test the snapserver
30 source registration).
31 """
32 media = MagicMock()
33 return SnapcastMAStream(
34 provider=provider,
35 media=media,
36 stream_name=name,
37 )
38
39
40@pytest.mark.parametrize(
41 ("sample_rate", "bit_depth", "expected"),
42 [
43 (48000, 16, "sampleformat=48000:16:2"),
44 (96000, 16, "sampleformat=96000:16:2"),
45 (96000, 24, "sampleformat=96000:24:2&packed_s24le=true"),
46 (192000, 24, "sampleformat=192000:24:2&packed_s24le=true"),
47 ],
48)
49def test_sampleformat_query_enables_packed_s24le_for_24bit(
50 sample_rate: int, bit_depth: int, expected: str
51) -> None:
52 """24-bit TCP sources must request packed_s24le; 16-bit must not."""
53 audio_format = snapcast_stream_format(sample_rate, bit_depth)
54 assert snapcast_sampleformat_query(audio_format) == expected
55
56
57def test_stream_format_maps_bit_depth_to_pcm_content_type() -> None:
58 """Bit depth selects the matching packed PCM content type."""
59 assert snapcast_stream_format(48000, 16).content_type == ContentType.PCM_S16LE
60 assert snapcast_stream_format(96000, 24).content_type == ContentType.PCM_S24LE
61
62
63def test_transport_format_uses_snapserver_codec() -> None:
64 """The reported final format follows the codec configured on Snapserver."""
65 provider = MagicMock()
66 provider._use_builtin_server = True
67 provider._snapcast_server_transport_codec = "opus"
68 provider.stream_audio_format = DEFAULT_SNAPCAST_FORMAT
69 stream = _make_stream(provider)
70
71 output_format = stream._get_transport_format()
72
73 assert output_format.content_type == ContentType.OPUS
74 assert output_format.codec_type == ContentType.OPUS
75 assert output_format.sample_rate == 48000
76 assert output_format.bit_depth == 16
77
78
79def test_transport_format_follows_configured_stream_format() -> None:
80 """Transport format inherits the configured Snapcast PCM sample rate/bit depth."""
81 provider = MagicMock()
82 provider._use_builtin_server = True
83 provider._snapcast_server_transport_codec = "flac"
84 provider.stream_audio_format = snapcast_stream_format(96000, 24)
85 stream = _make_stream(provider)
86
87 output_format = stream._get_transport_format()
88
89 assert output_format.content_type == ContentType.FLAC
90 assert output_format.sample_rate == 96000
91 assert output_format.bit_depth == 24
92
93
94def test_external_transport_format_reads_uri_codec() -> None:
95 """External Snapserver codec is read from the stream URI query."""
96 provider = MagicMock()
97 provider._use_builtin_server = False
98 provider.stream_audio_format = DEFAULT_SNAPCAST_FORMAT
99 stream = _make_stream(provider)
100 stream.snap_stream = MagicMock(
101 _stream={"uri": {"query": {"codec": "flac"}}},
102 )
103
104 output_format = stream._get_transport_format()
105
106 assert output_format.content_type == ContentType.FLAC
107 assert output_format.codec_type == ContentType.FLAC
108
109
110@pytest.mark.asyncio
111async def test_register_uses_configured_sampleformat(
112 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
113) -> None:
114 """stream_add_stream URI must include the configured sampleformat and packed_s24le."""
115 fake_provider.stream_audio_format = snapcast_stream_format(96000, 24)
116 fake_snapserver.queue_success(stream_id="hires-1")
117 stream = _make_stream(fake_provider)
118
119 await stream._register_tcp_server_source()
120
121 assert len(fake_snapserver.add_stream_calls) == 1
122 uri = fake_snapserver.add_stream_calls[0]
123 assert "sampleformat=96000:24:2" in uri
124 assert "packed_s24le=true" in uri
125
126
127def test_output_plan_is_registered_for_all_snapcast_members() -> None:
128 """Every client consuming a shared Snapcast stream gets the same output path."""
129 provider = MagicMock()
130 provider.stream_audio_format = DEFAULT_SNAPCAST_FORMAT
131 stream = _make_stream(provider)
132 stream.media.source_id = "queue-1"
133 stream.media.queue_session_id = "session-1"
134 stream.snap_stream = MagicMock(identifier="stream-1")
135 group = MagicMock(
136 stream="stream-1",
137 clients=["snap-child"],
138 )
139 provider._snapserver.groups = [group]
140 provider._get_ma_id.return_value = "child"
141 output_details = MagicMock(player_ids=["leader"])
142 stream._output_plan = MagicMock(output_details=output_details)
143
144 stream._register_output_plan()
145
146 registered_ids = {
147 call.args[0] for call in provider.mass.streams.audio_processing.update_output.call_args_list
148 }
149 assert registered_ids == {"leader", "child"}
150
151
152@pytest.mark.asyncio
153async def test_happy_path_register_succeeds_first_attempt(
154 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
155) -> None:
156 """Regression: a clean snapserver returns id on first try, MA registers the stream."""
157 fake_snapserver.queue_success(stream_id="ok-1")
158 stream = _make_stream(fake_provider)
159
160 await stream._register_tcp_server_source()
161
162 assert stream.snap_stream is not None
163 assert stream.snap_stream.identifier == "ok-1"
164 assert len(fake_snapserver.add_stream_calls) == 1
165 assert "sampleformat=48000:16:2" in fake_snapserver.add_stream_calls[0]
166 assert "packed_s24le" not in fake_snapserver.add_stream_calls[0]
167
168
169@pytest.mark.asyncio
170async def test_real_port_conflict_retries_with_different_port(
171 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
172) -> None:
173 """Regression: a non-name error (e.g. port already bound) keeps the loop going."""
174 fake_snapserver.queue_other_error("bind: Address already in use")
175 fake_snapserver.queue_success(stream_id="ok-after-retry")
176
177 stream = _make_stream(fake_provider)
178 await stream._register_tcp_server_source()
179
180 assert stream.snap_stream is not None
181 assert stream.snap_stream.identifier == "ok-after-retry"
182 # Two add_stream calls were made â first failed, second succeeded
183 assert len(fake_snapserver.add_stream_calls) == 2
184 # The two URIs must use different ports
185 port_1 = fake_snapserver.add_stream_calls[0].split("0.0.0.0:")[1].split("?")[0]
186 port_2 = fake_snapserver.add_stream_calls[1].split("0.0.0.0:")[1].split("?")[0]
187 assert port_1 != port_2
188
189
190@pytest.mark.asyncio
191async def test_all_attempts_exhausted_raises_with_honest_message(
192 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
193) -> None:
194 """
195 When all 50 retries fail the error must reference 'after retries' or similar.
196
197 Must NOT say 'No free port found' which lies about the cause.
198 """
199 # Queue 51 errors so even after 50 attempts the loop hits the raise
200 for _ in range(51):
201 fake_snapserver.queue_other_error("Some persistent error")
202
203 stream = _make_stream(fake_provider)
204
205 with pytest.raises(RuntimeError) as exc_info:
206 await stream._register_tcp_server_source()
207
208 message = str(exc_info.value)
209 assert "No free port" not in message, (
210 f"Error message still claims port shortage when the real cause may differ: {message}"
211 )
212 assert "attempts" in message.lower() or "register" in message.lower()
213
214
215@pytest.mark.parametrize(
216 ("result", "expected"),
217 [
218 (
219 {
220 "code": -32603,
221 "data": 'Stream with name "x" already exists',
222 "message": "Internal error",
223 },
224 True,
225 ),
226 (
227 {
228 "code": -32603,
229 "data": 'Stream with name "abc" already exists in registry',
230 "message": "Internal error",
231 },
232 True,
233 ),
234 # negative cases â these must NOT match
235 (
236 {"code": -32603, "data": "bind: Address already in use", "message": "Internal error"},
237 False,
238 ),
239 (
240 {
241 "code": -32603,
242 "data": "Some other error already exists somewhere",
243 "message": "Internal error",
244 },
245 False,
246 ),
247 ({}, False),
248 (None, False),
249 ("plain string", False),
250 ({"data": 12345}, False),
251 ],
252)
253def test_is_name_collision_error_matches_only_specific_pattern(
254 result: object, expected: bool
255) -> None:
256 """_is_name_collision_error must NOT false-positive on unrelated errors."""
257 assert SnapcastMAStream._is_name_collision_error(result) is expected
258
259
260@pytest.mark.asyncio
261async def test_name_collision_with_local_stream_cached_adopts_it(
262 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
263) -> None:
264 """
265 When snapserver reports the name as already-registered MA must adopt the orphan.
266
267 The stream must be in the local snapserver cache; MA must adopt it instead of
268 looping until retries exhaust.
269 """
270 target_name = "Music Assistant - 590b15 (announcement)"
271 fake_snapserver.cache_stream_directly("orphan-id", target_name)
272 fake_snapserver.queue_name_collision()
273
274 stream = _make_stream(fake_provider, name=target_name)
275 await stream._register_tcp_server_source()
276
277 assert stream.snap_stream is not None
278 assert stream.snap_stream.identifier == "orphan-id"
279 # Adoption must NOT spend further retries
280 assert len(fake_snapserver.add_stream_calls) == 1
281
282
283@pytest.mark.asyncio
284async def test_name_collision_with_incompatible_format_recreates_stream(
285 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
286) -> None:
287 """Orphan streams with a different sample format must be removed and recreated."""
288 target_name = "Music Assistant - hires"
289 fake_provider.stream_audio_format = snapcast_stream_format(96000, 24)
290 fake_snapserver.cache_stream_directly("orphan-id", target_name)
291 fake_snapserver.queue_name_collision()
292 fake_snapserver.queue_success(stream_id="recreated-id")
293
294 stream = _make_stream(fake_provider, name=target_name)
295 await stream._register_tcp_server_source()
296
297 assert stream.snap_stream is not None
298 assert stream.snap_stream.identifier == "recreated-id"
299 assert fake_snapserver.stream_remove_stream.await_count == 1
300 assert "sampleformat=96000:24:2" in fake_snapserver.add_stream_calls[-1]
301 assert "packed_s24le=true" in fake_snapserver.add_stream_calls[-1]
302
303
304@pytest.mark.asyncio
305async def test_orphan_stream_after_ma_restart_gets_adopted_via_resync(
306 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
307) -> None:
308 """
309 Bug B core scenario: MA restarted while a stream was registered.
310
311 The local snapserver cache is empty, but the server still holds the orphan.
312 MA must call status() + synchronize() to re-discover it before adopting.
313 """
314 target_name = "Music Assistant - 590b15 (announcement)"
315 # Stream visible only via status(), NOT in the local cache yet
316 fake_snapserver.stage_orphan_stream("orphan-id", target_name)
317 fake_snapserver.queue_name_collision()
318
319 # Sanity: the local cache is empty before
320 assert fake_snapserver._streams_by_id == {}
321
322 stream = _make_stream(fake_provider, name=target_name)
323 await stream._register_tcp_server_source()
324
325 assert stream.snap_stream is not None
326 assert stream.snap_stream.identifier == "orphan-id"
327 # The status() round-trip must have happened
328 assert fake_snapserver.status.await_count >= 1
329
330
331@pytest.mark.asyncio
332async def test_second_register_call_is_idempotent_when_stream_already_set(
333 fake_provider: MagicMock, fake_snapserver: FakeSnapserver
334) -> None:
335 """
336 Calling _register_tcp_server_source after self.snap_stream is set is a no-op.
337
338 This is the early-return guard at the top of the method. Real concurrency
339 safety is provided by `_lifecycle_lock` in setup(); that path is exercised
340 by integration tests, not unit tests.
341 """
342 fake_snapserver.queue_success(stream_id="single-stream")
343 fake_snapserver.queue_success(stream_id="should-not-happen")
344
345 stream = _make_stream(fake_provider)
346
347 await asyncio.gather(
348 stream._register_tcp_server_source(),
349 stream._register_tcp_server_source(),
350 )
351
352 assert stream.snap_stream is not None
353 assert stream.snap_stream.identifier == "single-stream"
354 # Only ONE add_stream call was issued
355 assert len(fake_snapserver.add_stream_calls) == 1
356
357
358def test_pick_port_avoids_already_tried() -> None:
359 """Within a single retry loop, _pick_port_avoiding must not return a tried port."""
360 # Build a stream stub that exposes only what _pick_port_avoiding needs
361 stream = SnapcastMAStream.__new__(SnapcastMAStream)
362
363 used: set[int] = set()
364 for _ in range(100):
365 port = stream._pick_port_avoiding(used)
366 assert port is not None, "Helper unexpectedly returned None when ports remain"
367 assert 4953 <= port <= 4953 + 200
368 assert port not in used
369 used.add(port)
370