Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 29 additions & 5 deletions livekit-rtc/livekit/rtc/audio_mixer.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,12 @@ def __init__(
"""
self._streams: set[_Stream] = set()
self._buffers: dict[_Stream, np.ndarray] = {}
# A read that timed out keeps running here and is awaited again on the
# next round. Cancelling it would also cancel the stream itself when the
# stream is an async generator, which then ends and drops all its audio.
self._pending: dict[_Stream, asyncio.Future[AudioFrame]] = {}
# Reads cancelled by remove_stream that may still be finishing their cancellation.
self._cancelling: set[asyncio.Future[AudioFrame]] = set()
self._sample_rate: int = sample_rate
self._num_channels: int = num_channels
self._chunk_size: int = blocksize if blocksize > 0 else int(sample_rate // 10)
Expand Down Expand Up @@ -90,6 +96,12 @@ def remove_stream(self, stream: AsyncIterator[AudioFrame]) -> None:
"""
self._streams.discard(stream)
self._buffers.pop(stream, None)
pending = self._pending.pop(stream, None)
if pending is not None:
pending.cancel()
# aclose() waits for the cancellation to finish, which can outlive this call.
self._cancelling.add(pending)
pending.add_done_callback(self._cancelling.discard)

def __aiter__(self) -> "AudioMixer":
return self
Expand All @@ -110,6 +122,14 @@ async def aclose(self) -> None:
self._mixer_task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await self._mixer_task
pending_reads = list(self._pending.values())
self._pending.clear()
for pending in pending_reads:
pending.cancel()
# Cancelling only requests cancellation; wait so stream cleanup has finished on return.
# Reads already cancelled by remove_stream are awaited but not cancelled a second time,
# which would interrupt their cleanup.
await asyncio.gather(*pending_reads, *self._cancelling, return_exceptions=True)

def end_input(self) -> None:
"""
Expand Down Expand Up @@ -174,13 +194,17 @@ async def _get_contribution(
had_data = buf.shape[0] > 0
exhausted = False
while buf.shape[0] < self._chunk_size and not exhausted:
try:
frame = await asyncio.wait_for(
stream.__anext__(), timeout=self._stream_timeout_ms / 1000
)
except asyncio.TimeoutError:
pending = self._pending.get(stream)
if pending is None:
pending = asyncio.ensure_future(stream.__anext__())
self._pending[stream] = pending
done, _ = await asyncio.wait({pending}, timeout=self._stream_timeout_ms / 1000)
if not done:
logger.warning(f"AudioMixer: stream {stream} timeout, ignoring")
break
self._pending.pop(stream, None)
try:
frame = pending.result()
except StopAsyncIteration:
exhausted = True
break
Expand Down
87 changes: 86 additions & 1 deletion tests/rtc/test_mixer.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
# type: ignore

import asyncio

import numpy as np
import pytest

from livekit.rtc import AudioMixer
from livekit.rtc import AudioFrame, AudioMixer
from livekit.rtc.utils import sine_wave_generator

SAMPLE_RATE = 48000
Expand Down Expand Up @@ -56,3 +58,86 @@ async def test_mixer_two_sine_waves():
# Assert that the peaks include 440Hz and 880Hz (with a tolerance of ±5 Hz)
assert any(np.isclose(peak_freqs, 440, atol=5)), f"Expected 440 Hz in peaks, got: {peak_freqs}"
assert any(np.isclose(peak_freqs, 880, atol=5)), f"Expected 880 Hz in peaks, got: {peak_freqs}"


@pytest.mark.asyncio
async def test_mixer_keeps_generator_stream_that_is_slower_than_the_timeout():
"""
A stream that misses the timeout once must not be lost: its audio is mixed
once it arrives.
"""

async def slow_stream():
await asyncio.sleep(0.3) # longer than stream_timeout_ms
for _ in range(5):
yield AudioFrame(
np.ones(BLOCKSIZE, dtype=np.int16).tobytes(), SAMPLE_RATE, 1, BLOCKSIZE
)

mixer = AudioMixer(sample_rate=SAMPLE_RATE, num_channels=1, stream_timeout_ms=100)
mixer.add_stream(slow_stream())
mixer.end_input()

async def read_all():
return [frame async for frame in mixer]

frames = await asyncio.wait_for(read_all(), timeout=5)
await mixer.aclose()

frames_with_audio = [f for f in frames if np.any(np.frombuffer(f.data.tobytes(), np.int16))]
assert len(frames_with_audio) == 5


@pytest.mark.asyncio
async def test_mixer_aclose_waits_for_pending_stream_cleanup():
"""
Closing the mixer must finish cancelling an in-flight stream read, so the
stream's cleanup has completed by the time `aclose` returns.
"""
cleaned_up = False

async def stream_with_slow_cleanup():
nonlocal cleaned_up
try:
await asyncio.sleep(10)
yield AudioFrame(
np.ones(BLOCKSIZE, dtype=np.int16).tobytes(), SAMPLE_RATE, 1, BLOCKSIZE
)
finally:
await asyncio.sleep(0.1)
cleaned_up = True

mixer = AudioMixer(sample_rate=SAMPLE_RATE, num_channels=1, stream_timeout_ms=50)
mixer.add_stream(stream_with_slow_cleanup())
await asyncio.sleep(0.2) # the read is pending and has missed the timeout

await mixer.aclose()

assert cleaned_up


@pytest.mark.asyncio
async def test_mixer_aclose_waits_for_cleanup_of_a_removed_stream():
"""A stream removed while its read is pending is also finished by `aclose`."""
cleaned_up = False

async def stream_with_slow_cleanup():
nonlocal cleaned_up
try:
await asyncio.sleep(10)
yield AudioFrame(
np.ones(BLOCKSIZE, dtype=np.int16).tobytes(), SAMPLE_RATE, 1, BLOCKSIZE
)
finally:
await asyncio.sleep(0.1)
cleaned_up = True

stream = stream_with_slow_cleanup()
mixer = AudioMixer(sample_rate=SAMPLE_RATE, num_channels=1, stream_timeout_ms=50)
mixer.add_stream(stream)
await asyncio.sleep(0.2)
mixer.remove_stream(stream)

await mixer.aclose()

assert cleaned_up
Loading