diff --git a/livekit-rtc/livekit/rtc/_utils.py b/livekit-rtc/livekit/rtc/_utils.py index 91175815..088fb0e1 100644 --- a/livekit-rtc/livekit/rtc/_utils.py +++ b/livekit-rtc/livekit/rtc/_utils.py @@ -113,6 +113,19 @@ async def wait_for(self, fnc: Callable[[T], bool]) -> T: self.task_done() + def release(self) -> None: + """Mark every unfinished item as done. + + Lets a ``join()`` that is already waiting on this queue return once there + is no consumer left for it. + """ + while True: + try: + self.task_done() + except ValueError: + # task_done() raises once every item is accounted for + return + class BroadcastQueue(Generic[T]): """Queue with multiple subscribers.""" @@ -135,6 +148,10 @@ def subscribe(self) -> Queue[T]: def unsubscribe(self, queue: Queue[T]) -> None: self._subscribers.remove(queue) + # A join() may already hold this queue. Events that are still queued, or that + # the subscriber took but never finished, have no consumer anymore, so + # release them instead of leaving the joiner waiting forever. + queue.release() async def join(self) -> None: async with self._lock: diff --git a/livekit-rtc/tests/test_broadcast_queue.py b/livekit-rtc/tests/test_broadcast_queue.py new file mode 100644 index 00000000..ff7b3745 --- /dev/null +++ b/livekit-rtc/tests/test_broadcast_queue.py @@ -0,0 +1,98 @@ +"""Unit tests for BroadcastQueue subscriber removal. + +``Room._listen_task`` calls ``BroadcastQueue.join()`` after every event, so a +subscriber that goes away (for example a cancelled ``publish_track``) must not +leave that join waiting on events nobody will finish. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from livekit.rtc._utils import BroadcastQueue + + +async def _join_finishes(queue: BroadcastQueue, timeout: float = 1.0) -> bool: + try: + await asyncio.wait_for(queue.join(), timeout) + except asyncio.TimeoutError: + return False + return True + + +@pytest.mark.asyncio +async def test_join_waits_for_a_subscriber_that_has_not_finished() -> None: + broadcast: BroadcastQueue[int] = BroadcastQueue() + subscriber = broadcast.subscribe() + broadcast.put_nowait(1) + + assert not await _join_finishes(broadcast, timeout=0.1) + + await subscriber.get() + subscriber.task_done() + assert await _join_finishes(broadcast) + + +@pytest.mark.asyncio +async def test_unsubscribe_releases_a_join_waiting_on_a_queued_event() -> None: + broadcast: BroadcastQueue[int] = BroadcastQueue() + subscriber = broadcast.subscribe() + broadcast.put_nowait(1) + + joiner = asyncio.create_task(broadcast.join()) + await asyncio.sleep(0) + assert not joiner.done() + + broadcast.unsubscribe(subscriber) + + await asyncio.wait_for(joiner, 1.0) + assert broadcast.len_subscribers() == 0 + + +@pytest.mark.asyncio +async def test_unsubscribe_releases_a_join_waiting_on_an_event_that_was_taken() -> None: + broadcast: BroadcastQueue[int] = BroadcastQueue() + subscriber = broadcast.subscribe() + broadcast.put_nowait(1) + await subscriber.get() # taken, but the subscriber never calls task_done() + + joiner = asyncio.create_task(broadcast.join()) + await asyncio.sleep(0) + assert not joiner.done() + + broadcast.unsubscribe(subscriber) + + await asyncio.wait_for(joiner, 1.0) + + +@pytest.mark.asyncio +async def test_unsubscribe_keeps_other_subscribers_in_the_join() -> None: + broadcast: BroadcastQueue[int] = BroadcastQueue() + gone = broadcast.subscribe() + staying = broadcast.subscribe() + broadcast.put_nowait(1) + + joiner = asyncio.create_task(broadcast.join()) + await asyncio.sleep(0) + broadcast.unsubscribe(gone) + await asyncio.sleep(0.05) + assert not joiner.done() + + await staying.get() + staying.task_done() + await asyncio.wait_for(joiner, 1.0) + + +@pytest.mark.asyncio +async def test_unsubscribe_after_every_event_is_finished_is_a_no_op() -> None: + broadcast: BroadcastQueue[int] = BroadcastQueue() + subscriber = broadcast.subscribe() + broadcast.put_nowait(1) + await subscriber.get() + subscriber.task_done() + + broadcast.unsubscribe(subscriber) + + assert await _join_finishes(broadcast)