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