/
/
1"""Tests for SharedGroupStream."""
2
3from __future__ import annotations
4
5import asyncio
6from collections.abc import AsyncIterator
7from typing import TYPE_CHECKING
8
9import pytest
10
11from music_assistant.providers.msx_bridge.provider import SharedGroupStream
12
13if TYPE_CHECKING:
14 from music_assistant.providers.msx_bridge.provider import MSXBridgeProvider
15
16
17async def _chunks(*data: bytes) -> AsyncIterator[bytes]:
18 """Yield bytes chunks as an async iterator."""
19 for chunk in data:
20 yield chunk
21
22
23async def _collect(stream: SharedGroupStream, player_id: str) -> list[bytes]:
24 """Subscribe and collect all chunks into a list."""
25 result = []
26 async for chunk in stream.subscribe(player_id):
27 result.append(chunk)
28 return result
29
30
31async def test_subscribe_receives_all_chunks() -> None:
32 """Subscriber should receive every chunk produced."""
33 stream = SharedGroupStream("g1", "uri://test")
34 await stream.start(_chunks(b"a", b"b", b"c"))
35 result = await _collect(stream, "tv1")
36 assert result == [b"a", b"b", b"c"]
37
38
39async def test_late_joiner_after_finish_does_not_hang() -> None:
40 """
41 subscribe() called after producer has already finished must not block indefinitely.
42
43 Regression test for: https://github.com/music-assistant/server/pull/3123#discussion_r2842897555
44 """
45 stream = SharedGroupStream("g1", "uri://test")
46 await stream.start(_chunks(b"x", b"y"))
47
48 # Wait for the producer to fully finish before subscribing.
49 assert stream.producer_task is not None
50 await asyncio.wait_for(stream.producer_task, timeout=5.0)
51 assert stream.finished is True
52
53 # Late subscriber: must get catch-up data and exit cleanly (no hang).
54 result = await asyncio.wait_for(_collect(stream, "late"), timeout=5.0)
55 assert result == [b"x", b"y"]
56
57
58async def test_late_joiner_with_empty_stream_does_not_hang() -> None:
59 """Late joiner on a stream that produced zero chunks must also exit cleanly."""
60 stream = SharedGroupStream("g1", "uri://test")
61 await stream.start(_chunks()) # no chunks
62
63 assert stream.producer_task is not None
64 await asyncio.wait_for(stream.producer_task, timeout=5.0)
65 assert stream.finished is True
66
67 result = await asyncio.wait_for(_collect(stream, "late"), timeout=5.0)
68 assert result == []
69
70
71async def test_concurrent_subscribers_receive_live_chunks() -> None:
72 """Multiple subscribers joining before stream starts all receive all chunks."""
73 stream = SharedGroupStream("g1", "uri://test")
74
75 async def slow_source() -> AsyncIterator[bytes]:
76 for i in range(3):
77 await asyncio.sleep(0)
78 yield bytes([i])
79
80 await stream.start(slow_source())
81
82 results = await asyncio.gather(
83 _collect(stream, "tv1"),
84 _collect(stream, "tv2"),
85 )
86 assert results[0] == results[1]
87 assert len(results[0]) == 3
88
89
90async def test_concurrent_replace_yields_single_stream(provider: MSXBridgeProvider) -> None:
91 """
92 Concurrent replacing get_or_create_shared_stream calls must yield ONE stream.
93
94 Without serialization, both callers pass the "existing" check while the old
95 producer is being awaited, each creates its own stream, and the loser's
96 ffmpeg producer is orphaned â consuming audio with zero subscribers.
97 """
98
99 async def infinite_source() -> AsyncIterator[bytes]:
100 while True:
101 await asyncio.sleep(0.01)
102 yield b"chunk"
103
104 old = await provider.get_or_create_shared_stream("g1", "uri://old", infinite_source())
105 try:
106 results = await asyncio.gather(
107 provider.get_or_create_shared_stream("g1", "uri://new", infinite_source()),
108 provider.get_or_create_shared_stream("g1", "uri://new", infinite_source()),
109 )
110 assert results[0] is results[1]
111 assert provider._shared_streams["g1"] is results[0]
112 finally:
113 await old.stop()
114 for stream in {id(s): s for s in provider._shared_streams.values()}.values():
115 await stream.stop()
116 await asyncio.gather(
117 *(s.stop() for s in results),
118 return_exceptions=True,
119 )
120
121
122async def test_cancel_stops_subscription() -> None:
123 """Cancelling a subscriber's task cleans up the subscriber registry."""
124
125 async def infinite_source() -> AsyncIterator[bytes]:
126 while True:
127 await asyncio.sleep(0.01)
128 yield b"chunk"
129
130 stream = SharedGroupStream("g1", "uri://test")
131 await stream.start(infinite_source())
132
133 task = asyncio.create_task(_collect(stream, "tv1"))
134 await asyncio.sleep(0.05) # let subscriber register and receive some chunks
135 task.cancel()
136 with pytest.raises(asyncio.CancelledError):
137 await task
138
139 # After cancel, subscriber should be cleaned up
140 assert "tv1" not in stream.subscribers
141
142 # Cleanup
143 if stream.producer_task:
144 stream.producer_task.cancel()
145 with pytest.raises(asyncio.CancelledError):
146 await stream.producer_task
147