diff --git a/CLAUDE.md b/CLAUDE.md index a2ae522..4b325fd 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -121,6 +121,7 @@ ts6-stream-bot/ │ │ ├── audio.py <- PulseAudio sink helpers (introspection) │ │ ├── audio_capture.py <- PulseAudio -> Opus -> TS3 voice frames │ │ ├── video_capture.py <- x11grab + Pulse via aiortc MediaPlayers +│ │ ├── video_broadcaster.py <- Single libvpx encoder, per-viewer av.Packet fan-out │ │ ├── stream_signaling.py <- TS6 stream signaling (setupstream etc.) │ │ └── stream_publisher.py <- Per-viewer aiortc RTCPeerConnection │ │ diff --git a/src/ts6_stream_bot/pipeline/controller.py b/src/ts6_stream_bot/pipeline/controller.py index c07d616..8621b81 100644 --- a/src/ts6_stream_bot/pipeline/controller.py +++ b/src/ts6_stream_bot/pipeline/controller.py @@ -40,6 +40,10 @@ from ts6_stream_bot.pipeline.browser import BrowserManager from ts6_stream_bot.pipeline.stream_publisher import StreamPublisher from ts6_stream_bot.pipeline.stream_signaling import StreamSignaling +from ts6_stream_bot.pipeline.video_broadcaster import ( + VideoBroadcaster, + VideoBroadcasterConfig, +) from ts6_stream_bot.pipeline.video_capture import VideoCapture, VideoCaptureConfig from ts6_stream_bot.sources import StreamSource, resolve_source from ts6_stream_bot.ts3lib.client import Ts3Client, Ts3ClientOptions @@ -254,7 +258,25 @@ async def _allocate_stream(self) -> None: pulse_source=f"{settings.PULSE_SINK}.monitor", ) capture = VideoCapture(capture_config) - publisher = StreamPublisher(client=self._ts3_client, signaling=signaling, capture=capture) + # Single libvpx instance shared across viewers. STREAM_BITRATE is + # in kbps (mirrors the TS6 setupstream parameter); the codec + # context wants bps. The factory closure defers reading + # capture.video_track until publisher.start() has booted capture. + broadcaster = VideoBroadcaster( + source_track_factory=lambda: capture.video_track, + config=VideoBroadcasterConfig( + bitrate=settings.STREAM_BITRATE * 1000, + width=settings.SCREEN_WIDTH, + height=settings.SCREEN_HEIGHT, + framerate=settings.SCREEN_FPS, + ), + ) + publisher = StreamPublisher( + client=self._ts3_client, + signaling=signaling, + capture=capture, + video_broadcaster=broadcaster, + ) log.info( "controller.stream_setup", diff --git a/src/ts6_stream_bot/pipeline/stream_publisher.py b/src/ts6_stream_bot/pipeline/stream_publisher.py index 393b827..0b86bce 100644 --- a/src/ts6_stream_bot/pipeline/stream_publisher.py +++ b/src/ts6_stream_bot/pipeline/stream_publisher.py @@ -43,6 +43,7 @@ SignalingType, StreamSignaling, ) +from ts6_stream_bot.pipeline.video_broadcaster import VideoBroadcaster from ts6_stream_bot.pipeline.video_capture import VideoCapture from ts6_stream_bot.ts3lib.client import Ts3Client @@ -78,10 +79,12 @@ def __init__( client: Ts3Client, signaling: StreamSignaling, capture: VideoCapture, + video_broadcaster: VideoBroadcaster, ) -> None: self._client = client self._signaling = signaling self._capture = capture + self._video_broadcaster = video_broadcaster self._stream_id: str | None = None self._stream_started_event = asyncio.Event() @@ -90,9 +93,11 @@ def __init__( self._lock = asyncio.Lock() # Track in-flight tasks so the GC doesn't clean them up mid-flight. self._tasks: set[asyncio.Task[None]] = set() - # MediaRelay fans out one source track to many per-viewer subscribers - # so two PeerConnections don't both call recv() on the same underlying - # track at the same time (would race on the parec / x11grab subprocess). + # MediaRelay still handles audio fan-out: one parec subprocess + # delivers raw PCM and we want each viewer's RTCRtpSender to read + # its own subscriber queue rather than race on the source. Video + # bypasses the relay entirely - it goes through the broadcaster + # which encodes once and ships pre-encoded packets per viewer. self._relay = MediaRelay() # Wire up signaling callbacks. We chain so existing handlers stay alive. @@ -144,8 +149,10 @@ async def start( return self._stream_id # Make sure the underlying ffmpeg pipelines are running before we - # accept any join requests. + # accept any join requests. Capture has to come up first so the + # broadcaster can resolve the source track factory. await self._capture.start() + await self._video_broadcaster.start() self._stream_started_event.clear() self._signaling.send_setup_stream( @@ -160,6 +167,7 @@ async def start( try: await asyncio.wait_for(self._stream_started_event.wait(), timeout=timeout) except TimeoutError: + await self._video_broadcaster.stop() await self._capture.stop() raise @@ -170,6 +178,7 @@ async def start( async def stop(self) -> None: """Kick all viewers, stopstream, and tear down capture.""" if self._stream_id is None: + await self._video_broadcaster.stop() await self._capture.stop() return @@ -193,6 +202,7 @@ async def stop(self) -> None: # we tear down ffmpeg out from under the encoder. await asyncio.sleep(0.5) + await self._video_broadcaster.stop() await self._capture.stop() self._stream_id = None log.info("stream_publisher.stopped", stream_id=stream_id) @@ -314,12 +324,16 @@ async def _on_ice_state_change() -> None: ice=pc.iceConnectionState, ) - # Each viewer needs its OWN track that pulls from the shared source - # via MediaRelay. Sharing the source track directly across PCs makes - # both senders call recv() concurrently and crashes parec's - # readexactly() with a "another coroutine is already waiting" error. - if self._capture.video_track is not None: - pc.addTrack(self._relay.subscribe(self._capture.video_track)) + # Video: subscribe to the broadcaster's pre-encoded fan-out. The + # returned track yields ``av.Packet`` objects, which RTCRtpSender + # routes through ``encoder.pack()`` (RTP packetization only, no + # libvpx call) - so each viewer's encoder spend is microseconds + # rather than the ~750 MB per-viewer libvpx setup we used to pay + # with MediaRelay-on-raw-frames. + pc.addTrack(self._video_broadcaster.subscribe()) + # Audio still goes through MediaRelay: parec produces raw PCM and + # each viewer's RTCRtpSender encodes Opus on its own. Opus is + # cheap enough that the broadcaster pattern doesn't pay back here. if self._capture.audio_track is not None: pc.addTrack(self._relay.subscribe(self._capture.audio_track)) diff --git a/src/ts6_stream_bot/pipeline/video_broadcaster.py b/src/ts6_stream_bot/pipeline/video_broadcaster.py new file mode 100644 index 0000000..4e8ba61 --- /dev/null +++ b/src/ts6_stream_bot/pipeline/video_broadcaster.py @@ -0,0 +1,319 @@ +"""Single-encoder video broadcast. + +Replaces the MediaRelay-based fan-out for the video track. The earlier +design subscribed every viewer's ``RTCRtpSender`` to a copy of the raw +x11grab frames, and each sender ran its own libvpx encoder; per-viewer +RSS measured ~750 MB at 720p30 in live tests, which is what made the +4 GB host with 5+ viewers a non-starter. + +Here a single ``av.CodecContext`` runs in a pump task, encodes each +captured frame exactly once, and fans out the resulting ``av.Packet`` +to per-viewer ``asyncio.Queue``s. Each subscriber exposes a +``BroadcastVideoTrack`` whose ``recv()`` returns ``av.Packet`` (not +``av.VideoFrame``); aiortc's ``RTCRtpSender`` then drops into the +pre-encoded path:: + + if isinstance(data, Frame): + ... encode ... + else: + payloads, timestamp = self.__encoder.pack(data) # cheap RTP packetize + +So the only per-viewer cost is the RTP packetization plus the SRTP +output - libvpx itself runs once for the whole channel. + +Trade-offs we accept for the win: + +* Shared encoder = shared bitrate. We don't honour per-viewer REMB + hints; the codec runs at the configured ``STREAM_BITRATE``. The + earlier design didn't actually use REMB-driven adaptation either + (we set static bitrate) so this is a no-op in practice. +* Any one viewer's PLI / FIR causes a keyframe for all viewers + rather than just that one. Acceptable: keyframes are infrequent + (gop_size=3000 frames = 100 s at 30 fps) and a few extra ones + cost less than running N encoders. +* Slow viewers get frames dropped from their personal queue rather + than holding up everyone else. The connection-state callback in + StreamPublisher already evicts persistently broken peers. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import multiprocessing +from collections.abc import Callable +from dataclasses import dataclass + +import av +import structlog +from aiortc.mediastreams import MediaStreamError, MediaStreamTrack +from av.packet import Packet +from av.video.codeccontext import VideoCodecContext +from av.video.frame import PictureType, VideoFrame + +log = structlog.get_logger(__name__) + + +@dataclass(slots=True) +class VideoBroadcasterConfig: + """Encoder + queue knobs. Defaults mirror aiortc's per-sender + ``Vp8Encoder`` so the wire-level behaviour is unchanged - only + the number of times we run the encoder differs.""" + + bitrate: int # bits per second (note: STREAM_BITRATE in .env is kbps) + width: int + height: int + framerate: int + cpu_used: int = -6 + deadline: str = "realtime" + gop_size: int = 3000 + qmin: int = 2 + qmax: int = 56 + queue_size: int = 30 # ~1 second at 30 fps; slow viewers drop oldest + thread_count: int | None = None # None = auto-tune via aiortc's heuristic + + +class BroadcastVideoTrack(MediaStreamTrack): + """Per-viewer track that yields pre-encoded ``av.Packet``s. + + aiortc's ``RTCRtpSender`` checks ``isinstance(data, Frame)`` - + since we return ``Packet``, it calls ``encoder.pack(data)`` for + cheap RTP packetization instead of ``encoder.encode(frame)``. + """ + + kind = "video" + + def __init__( + self, + broadcaster: VideoBroadcaster, + queue: asyncio.Queue[Packet | None], + ) -> None: + super().__init__() + self._broadcaster = broadcaster + self._queue = queue + + async def recv(self) -> Packet: + packet = await self._queue.get() + if packet is None: + # Sentinel: broadcaster shutting down. Surface as the + # MediaStreamError aiortc expects to clean up the sender. + raise MediaStreamError + return packet + + def stop(self) -> None: + super().stop() + self._broadcaster._unsubscribe(self) + + +@dataclass(slots=True) +class _Subscriber: + queue: asyncio.Queue[Packet | None] + track: BroadcastVideoTrack + drops: int = 0 + + +class VideoBroadcaster: + """One libvpx encoder. Many viewer queues.""" + + def __init__( + self, + source_track_factory: Callable[[], MediaStreamTrack | None], + config: VideoBroadcasterConfig, + ) -> None: + # The factory indirection lets us be constructed before VideoCapture + # has started (capture.video_track is None until capture.start()). + # ``start()`` resolves the factory and only then commits to a track. + self._source_factory = source_track_factory + self._config = config + + self._codec: VideoCodecContext | None = None + self._source: MediaStreamTrack | None = None + self._pump_task: asyncio.Task[None] | None = None + self._stopped = False + self._force_keyframe = False + self._subscribers: list[_Subscriber] = [] + self._encoded_frames = 0 + + # --- lifecycle -------------------------------------------------------- + + async def start(self) -> None: + if self._pump_task is not None: + return + track = self._source_factory() + if track is None: + raise RuntimeError("video broadcaster: source track is None at start") + self._source = track + self._stopped = False + self._pump_task = asyncio.create_task(self._pump_loop(), name="video-broadcaster-pump") + log.info( + "video_broadcaster.started", + width=self._config.width, + height=self._config.height, + framerate=self._config.framerate, + bitrate=self._config.bitrate, + ) + + async def stop(self) -> None: + self._stopped = True + # Wake every subscriber out of recv() so the per-viewer sender + # can shut down cleanly. + for sub in list(self._subscribers): + with contextlib.suppress(asyncio.QueueFull): + sub.queue.put_nowait(None) + if self._pump_task is not None: + self._pump_task.cancel() + with contextlib.suppress(asyncio.CancelledError, Exception): + await self._pump_task + self._pump_task = None + # Flush any trailing encoded packets so libvpx releases internal buffers. + if self._codec is not None: + with contextlib.suppress(Exception): + for _ in self._codec.encode(None): + pass + self._codec = None + log.info("video_broadcaster.stopped", encoded_frames=self._encoded_frames) + + # --- subscription ----------------------------------------------------- + + def subscribe(self) -> BroadcastVideoTrack: + """Return a new track. Forces a keyframe so the joining viewer + can start decoding without waiting for the next periodic I-frame.""" + queue: asyncio.Queue[Packet | None] = asyncio.Queue(maxsize=self._config.queue_size) + track = BroadcastVideoTrack(self, queue) + sub = _Subscriber(queue=queue, track=track) + self._subscribers.append(sub) + self._force_keyframe = True + log.info( + "video_broadcaster.subscribed", + total_subscribers=len(self._subscribers), + ) + return track + + def _unsubscribe(self, track: BroadcastVideoTrack) -> None: + before = len(self._subscribers) + self._subscribers = [s for s in self._subscribers if s.track is not track] + after = len(self._subscribers) + if before != after: + log.info( + "video_broadcaster.unsubscribed", + total_subscribers=after, + ) + + # --- internals -------------------------------------------------------- + + def _build_codec(self, frame: VideoFrame) -> VideoCodecContext: + cfg = self._config + # Defaults below mirror aiortc.codecs.vpx.Vp8Encoder so the + # wire-level encoder behaviour is unchanged. PyAV's stubs return + # ``CodecContext`` from ``create()``; for a video codec the + # actual instance is a ``VideoCodecContext`` with the width / + # height / pix_fmt knobs we need. + codec: VideoCodecContext = av.CodecContext.create("libvpx", "w") + codec.width = frame.width + codec.height = frame.height + codec.bit_rate = cfg.bitrate + codec.pix_fmt = "yuv420p" + codec.gop_size = cfg.gop_size + codec.qmin = cfg.qmin + codec.qmax = cfg.qmax + codec.options = { + "bufsize": str(cfg.bitrate), + "cpu-used": str(cfg.cpu_used), + "deadline": cfg.deadline, + "lag-in-frames": "0", + "minrate": str(cfg.bitrate), + "maxrate": str(cfg.bitrate), + "noise-sensitivity": "4", + "overshoot-pct": "15", + "partitions": "0", + "static-thresh": "1", + "undershoot-pct": "100", + } + codec.thread_count = cfg.thread_count or _auto_thread_count( + frame.width * frame.height, multiprocessing.cpu_count() + ) + return codec + + async def _pump_loop(self) -> None: + try: + await self._pump() + except asyncio.CancelledError: + raise + except Exception as exc: + log.exception("video_broadcaster.pump_crashed", error=str(exc)) + # Wake subscribers so their senders fail fast rather than hang. + for sub in list(self._subscribers): + with contextlib.suppress(asyncio.QueueFull): + sub.queue.put_nowait(None) + + async def _pump(self) -> None: + assert self._source is not None + while not self._stopped: + try: + frame = await self._source.recv() + except MediaStreamError: + log.info("video_broadcaster.source_ended") + return + + if not isinstance(frame, VideoFrame): + # Source drift - shouldn't happen with x11grab, log and skip. + log.warning("video_broadcaster.unexpected_frame_type", got=type(frame).__name__) + continue + + if frame.format.name != "yuv420p": + frame = frame.reformat(format="yuv420p") + + if self._codec is None: + self._codec = self._build_codec(frame) + + if self._force_keyframe: + frame.pict_type = PictureType.I + self._force_keyframe = False + + try: + packets = list(self._codec.encode(frame)) + except av.error.FFmpegError as exc: + log.warning("video_broadcaster.encode_failed", error=str(exc)) + continue + + for packet in packets: + self._encoded_frames += 1 + self._fanout(packet) + + def _fanout(self, packet: Packet) -> None: + """Push a packet into every subscriber queue. On QueueFull, drop + the oldest packet for that subscriber rather than block - one + slow viewer must not stall the others.""" + for sub in list(self._subscribers): + try: + sub.queue.put_nowait(packet) + continue + except asyncio.QueueFull: + pass + # Make room and retry once. If the second put still fails, + # we silently skip this packet for this viewer; the encoder + # will produce another and the queue is bounded so this + # bounded-skip is the only way out. + with contextlib.suppress(asyncio.QueueEmpty): + sub.queue.get_nowait() + sub.drops += 1 + with contextlib.suppress(asyncio.QueueFull): + sub.queue.put_nowait(packet) + + +def _auto_thread_count(pixels: int, cpu_count: int) -> int: + """Match aiortc's libvpx thread heuristic (vpx.number_of_threads).""" + if pixels >= 1920 * 1080: + return min(cpu_count, 8) + if pixels >= 1280 * 720: + return min(cpu_count, 4) + if pixels >= 640 * 480: + return min(cpu_count, 2) + return 1 + + +__all__ = [ + "BroadcastVideoTrack", + "VideoBroadcaster", + "VideoBroadcasterConfig", +] diff --git a/tests/test_stream_publisher.py b/tests/test_stream_publisher.py index 5269058..ea63be8 100644 --- a/tests/test_stream_publisher.py +++ b/tests/test_stream_publisher.py @@ -114,11 +114,37 @@ def _emit(client: _FakeClient, raw: str) -> None: # --- fixture: publisher with mocked aiortc / capture -------------------- +class _FakeBroadcaster: + """Stand-in for VideoBroadcaster so the publisher tests don't need + a real libvpx context. ``subscribe()`` returns a fresh sentinel + object per call so each viewer's ``addTrack`` argument is unique + and the multi-viewer regression test can count subscriptions.""" + + def __init__(self) -> None: + self.started = False + self.stopped = False + self.subscribe_calls = 0 + self.tracks: list[Any] = [] + + async def start(self) -> None: + self.started = True + + async def stop(self) -> None: + self.stopped = True + + def subscribe(self) -> Any: + self.subscribe_calls += 1 + track = MagicMock(name=f"broadcast-track-{self.subscribe_calls}") + self.tracks.append(track) + return track + + @pytest.fixture def wired(monkeypatch): client = _FakeClient.make() sig = StreamSignaling(client) # type: ignore[arg-type] capture = _FakeCapture() + broadcaster = _FakeBroadcaster() pc_mocks: list[MagicMock] = [] @@ -131,13 +157,18 @@ def _factory(*args: Any, **kwargs: Any) -> MagicMock: # publisher's `RTCPeerConnection()` calls return the mock. monkeypatch.setattr("ts6_stream_bot.pipeline.stream_publisher.RTCPeerConnection", _factory) - publisher = StreamPublisher(client=client, signaling=sig, capture=capture) # type: ignore[arg-type] + publisher = StreamPublisher( # type: ignore[arg-type] + client=client, + signaling=sig, + capture=capture, + video_broadcaster=broadcaster, + ) # Stub the MediaRelay so .subscribe() returns the source track verbatim. - # That way the existing assertions on pc.addTrack() still work - # while we still get to verify subscribe() was actually called. + # Audio still flows through the relay; video bypasses it via the + # broadcaster fake above. publisher._relay = MagicMock() publisher._relay.subscribe = MagicMock(side_effect=lambda track: track) - return publisher, client, sig, capture, pc_mocks + return publisher, client, sig, capture, broadcaster, pc_mocks # --- start --------------------------------------------------------------- @@ -145,7 +176,7 @@ def _factory(*args: Any, **kwargs: Any) -> MagicMock: @pytest.mark.asyncio async def test_start_sends_setupstream_and_waits_for_started_notification(wired) -> None: - publisher, client, sig, capture, _ = wired + publisher, client, sig, capture, _bc, _ = wired async def _trigger_started() -> None: await asyncio.sleep(0.01) @@ -179,7 +210,7 @@ async def test_start_ignores_started_for_other_clid(wired) -> None: """If the server emits notifystreamstarted for someone else's stream (shouldn't happen in practice, but be defensive), our start() must not latch onto it.""" - publisher, _client, sig, _capture, _ = wired + publisher, _client, sig, _capture, _bc, _ = wired async def _trigger_other_then_ours() -> None: await asyncio.sleep(0.01) @@ -221,7 +252,7 @@ async def _trigger_other_then_ours() -> None: @pytest.mark.asyncio async def test_join_request_creates_pc_and_sends_offer(wired) -> None: - publisher, client, _sig, capture, pcs = wired + publisher, client, _sig, capture, _bc, pcs = wired # Skip start(); set the stream id directly. publisher._stream_id = "s1" @@ -236,11 +267,14 @@ async def test_join_request_creates_pc_and_sends_offer(wired) -> None: assert len(pcs) == 1 pc = pcs[0] - pc.addTrack.assert_any_call(capture.video_track) + # Video now comes from the broadcaster (one libvpx instance for the + # whole stream); audio still flows through the per-PC MediaRelay + # subscription so each viewer's Opus encoder reads its own queue + # rather than racing on the parec subprocess. + assert _bc.subscribe_calls == 1 + pc.addTrack.assert_any_call(_bc.tracks[0]) pc.addTrack.assert_any_call(capture.audio_track) - # Each of the two tracks must go through the relay so multiple PCs don't - # race on the underlying parec / x11grab subprocess. - assert publisher._relay.subscribe.call_count == 2 # type: ignore[attr-defined] + assert publisher._relay.subscribe.call_count == 1 # type: ignore[attr-defined] join_resp = next( parse_command(c) for c in client.sent if c.startswith("respondjoinstreamrequest") @@ -252,13 +286,15 @@ async def test_join_request_creates_pc_and_sends_offer(wired) -> None: @pytest.mark.asyncio -async def test_two_viewers_each_get_their_own_relay_subscription(wired) -> None: +async def test_two_viewers_each_get_their_own_track_subscription(wired) -> None: """Multi-viewer regression: a second join must NOT reuse the first - viewer's track - both go through MediaRelay.subscribe so each PC has - its own consumer queue. This is what was crashing the live deploy - with 'readexactly() called while another coroutine is already - waiting for incoming data'.""" - publisher, client, _sig, _capture, pcs = wired + viewer's track. Video gets one broadcaster.subscribe() per viewer + (each viewer reads its own per-subscriber queue of av.Packets); + audio stays with MediaRelay.subscribe() so each Opus encoder reads + its own queue rather than racing on the shared parec recv(). + This is what was crashing the live deploy with 'readexactly() + called while another coroutine is already waiting for incoming data'.""" + publisher, client, _sig, _capture, _bc, pcs = wired publisher._stream_id = "s1" _emit(client, "notifyjoinstreamrequest id=s1 clid=42") @@ -269,8 +305,10 @@ async def test_two_viewers_each_get_their_own_relay_subscription(wired) -> None: break assert len(pcs) == 2, f"expected 2 PCs, got {len(pcs)}" - # 2 viewers * 2 tracks each = 4 subscribe calls - assert publisher._relay.subscribe.call_count == 4 # type: ignore[attr-defined] + # One broadcaster subscription per viewer, plus one relay subscription + # per viewer for audio. + assert _bc.subscribe_calls == 2 + assert publisher._relay.subscribe.call_count == 2 # type: ignore[attr-defined] # --- answer flow --------------------------------------------------------- @@ -278,7 +316,7 @@ async def test_two_viewers_each_get_their_own_relay_subscription(wired) -> None: @pytest.mark.asyncio async def test_answer_calls_set_remote_description(wired) -> None: - publisher, client, sig, _capture, pcs = wired + publisher, client, sig, _capture, _bc, pcs = wired publisher._stream_id = "s1" _emit(client, "notifyjoinstreamrequest id=s1 clid=42") @@ -307,7 +345,7 @@ async def test_answer_calls_set_remote_description(wired) -> None: @pytest.mark.asyncio async def test_ice_candidate_is_forwarded_to_pc(wired, monkeypatch) -> None: - publisher, client, sig, _capture, pcs = wired + publisher, client, sig, _capture, _bc, pcs = wired publisher._stream_id = "s1" # Patch candidate_from_sdp so we don't have to construct a real @@ -361,7 +399,7 @@ async def test_ice_candidate_arriving_before_answer_is_buffered(wired, monkeypat """Regression for the live deploy where every ICE candidate arrived before the answer task finished and aiortc dropped the call. Now candidates wait for the answer via the per-viewer remote_set event.""" - publisher, client, sig, _capture, pcs = wired + publisher, client, sig, _capture, _bc, pcs = wired publisher._stream_id = "s1" fake_candidate = MagicMock() monkeypatch.setattr( @@ -408,7 +446,7 @@ async def test_ice_candidate_arriving_before_answer_is_buffered(wired, monkeypat @pytest.mark.asyncio async def test_client_left_closes_pc_and_drops_viewer(wired) -> None: - publisher, client, _sig, _capture, pcs = wired + publisher, client, _sig, _capture, _bc, pcs = wired publisher._stream_id = "s1" _emit(client, "notifyjoinstreamrequest id=s1 clid=42") @@ -429,7 +467,7 @@ async def test_client_left_closes_pc_and_drops_viewer(wired) -> None: @pytest.mark.asyncio async def test_stop_kicks_all_viewers_and_sends_stopstream(wired) -> None: - publisher, client, _sig, capture, pcs = wired + publisher, client, _sig, capture, _bc, pcs = wired publisher._stream_id = "s1" _emit(client, "notifyjoinstreamrequest id=s1 clid=42") @@ -453,7 +491,7 @@ async def test_stop_kicks_all_viewers_and_sends_stopstream(wired) -> None: @pytest.mark.asyncio async def test_stop_without_start_only_stops_capture(wired) -> None: - publisher, _client, _sig, capture, _ = wired + publisher, _client, _sig, capture, _bc, _ = wired await publisher.stop() assert capture.stopped is True assert publisher._stream_id is None @@ -464,7 +502,7 @@ async def test_stop_without_start_only_stops_capture(wired) -> None: @pytest.mark.asyncio async def test_status_reflects_state(wired) -> None: - publisher, _, _, _, _pcs = wired + publisher, _, _, _, _bc, _pcs = wired s = publisher.status() assert s.streaming is False assert s.viewer_count == 0 diff --git a/tests/test_video_broadcaster.py b/tests/test_video_broadcaster.py new file mode 100644 index 0000000..3d8c70b --- /dev/null +++ b/tests/test_video_broadcaster.py @@ -0,0 +1,217 @@ +"""Tests for the single-encoder video broadcaster. + +The broadcaster's job is to encode raw frames once and fan the resulting +``av.Packet``s out to N per-viewer queues. The expensive part (libvpx) +runs in a real test below to catch breakage; the rest of the surface +(subscription, drop-on-slow-consumer, sentinel-on-shutdown) is checked +with synthetic packets so the unit tests stay fast. +""" + +from __future__ import annotations + +import asyncio +from fractions import Fraction +from typing import Any +from unittest.mock import MagicMock + +import numpy as np +import pytest +from aiortc.mediastreams import MediaStreamError, MediaStreamTrack +from av.video.frame import VideoFrame + +from ts6_stream_bot.pipeline.video_broadcaster import ( + BroadcastVideoTrack, + VideoBroadcaster, + VideoBroadcasterConfig, +) + + +def _make_config(**overrides: Any) -> VideoBroadcasterConfig: + base = { + "bitrate": 500_000, + "width": 320, + "height": 240, + "framerate": 15, + "queue_size": 4, + } + base.update(overrides) + return VideoBroadcasterConfig(**base) + + +def _synthetic_frame(width: int = 320, height: int = 240, pts: int = 0) -> VideoFrame: + """Build a tiny YUV420p frame in pure numpy. Avoids decoding a real + video file from disk - tests stay self-contained.""" + arr = np.zeros((height * 3 // 2, width), dtype=np.uint8) + frame = VideoFrame.from_ndarray(arr, format="yuv420p") + frame.pts = pts + frame.time_base = Fraction(1, 90000) + return frame + + +class _FrameSource(MediaStreamTrack): + """Test track that emits a fixed number of synthetic frames then + raises MediaStreamError. Mirrors how the real x11grab MediaPlayer + behaves at end-of-input.""" + + kind = "video" + + def __init__(self, count: int, *, width: int = 320, height: int = 240) -> None: + super().__init__() + self._count = count + self._emitted = 0 + self._width = width + self._height = height + + async def recv(self) -> VideoFrame: + if self._emitted >= self._count: + raise MediaStreamError + frame = _synthetic_frame(self._width, self._height, pts=self._emitted * 6000) + self._emitted += 1 + return frame + + +# --- subscription mechanics ---------------------------------------------- + + +async def test_subscribe_returns_track_with_video_kind() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + track = bc.subscribe() + assert isinstance(track, BroadcastVideoTrack) + assert track.kind == "video" + + +async def test_subscribe_forces_keyframe_flag() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + bc._force_keyframe = False + bc.subscribe() + assert bc._force_keyframe is True + + +async def test_unsubscribe_removes_subscriber() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + track = bc.subscribe() + assert len(bc._subscribers) == 1 + track.stop() + assert len(bc._subscribers) == 0 + + +# --- recv() / sentinel --------------------------------------------------- + + +async def test_recv_returns_queued_packet() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + track = bc.subscribe() + sentinel = MagicMock(name="packet") + bc._subscribers[0].queue.put_nowait(sentinel) + got = await track.recv() + assert got is sentinel + + +async def test_recv_raises_media_stream_error_on_none_sentinel() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + track = bc.subscribe() + bc._subscribers[0].queue.put_nowait(None) + with pytest.raises(MediaStreamError): + await track.recv() + + +# --- fanout / slow-consumer behaviour ------------------------------------ + + +def test_fanout_distributes_packet_to_all_subscribers() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + a = bc.subscribe() + b = bc.subscribe() + pkt = MagicMock(name="packet") + + bc._fanout(pkt) + + assert bc._subscribers[0].queue.get_nowait() is pkt + assert bc._subscribers[1].queue.get_nowait() is pkt + # Sanity: tracks are distinct, queues are distinct. + assert a is not b + + +def test_full_queue_drops_oldest_and_increments_drop_counter() -> None: + """One slow viewer must not stall the rest. We drop the oldest packet + in their personal queue when it's full.""" + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config(queue_size=2)) + bc.subscribe() + + p0, p1, p2 = (MagicMock(name=f"p{i}") for i in range(3)) + + bc._fanout(p0) + bc._fanout(p1) + bc._fanout(p2) # queue is full when this lands + + sub = bc._subscribers[0] + # Oldest (p0) was evicted to make room for p2; p1 is still there. + remaining = [sub.queue.get_nowait() for _ in range(sub.queue.qsize())] + assert remaining == [p1, p2] + assert sub.drops == 1 + + +def test_one_slow_subscriber_doesnt_starve_a_fast_one() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config(queue_size=1)) + bc.subscribe() # slow viewer (we'll never drain) + bc.subscribe() # fast viewer (we'll drain each tick) + + fast_packets: list[MagicMock] = [] + for i in range(5): + pkt = MagicMock(name=f"pkt{i}") + bc._fanout(pkt) + fast_packets.append(bc._subscribers[1].queue.get_nowait()) + + assert fast_packets == [m for m in fast_packets] # 5 deliveries + assert len(fast_packets) == 5 + # Slow viewer's queue still has the latest, and it's exactly the bound. + assert bc._subscribers[0].queue.qsize() == 1 + + +# --- end-to-end pump (real libvpx) --------------------------------------- + + +async def test_pump_real_encoder_produces_packets_for_subscribers() -> None: + """Smoke test against actual libvpx via PyAV. Doesn't validate the + bytes - we just confirm the full path encode -> fanout -> recv() + delivers a Packet to every subscriber. If this regresses we've + broken pre-encoded mode and aiortc's RTCRtpSender will fall back + to wanting frames.""" + bc = VideoBroadcaster(lambda: _FrameSource(10), _make_config()) + track_a = bc.subscribe() + track_b = bc.subscribe() + + await bc.start() + + async def first_packet(track: BroadcastVideoTrack) -> Any: + return await asyncio.wait_for(track.recv(), timeout=5.0) + + pkt_a, pkt_b = await asyncio.gather(first_packet(track_a), first_packet(track_b)) + + # Real av.Packet, with bytes available - that's what RTCRtpSender.pack expects. + assert pkt_a is not None + assert pkt_b is not None + assert len(bytes(pkt_a)) > 0 + assert len(bytes(pkt_b)) > 0 + + await bc.stop() + + +async def test_stop_wakes_pending_recv_with_sentinel() -> None: + bc = VideoBroadcaster(lambda: _FrameSource(0), _make_config()) + track = bc.subscribe() + await bc.start() + + # No frames will arrive (source is empty); stop should wake recv(). + recv_task = asyncio.create_task(track.recv()) + await asyncio.sleep(0.05) + await bc.stop() + + with pytest.raises(MediaStreamError): + await asyncio.wait_for(recv_task, timeout=2.0) + + +async def test_start_raises_when_factory_returns_none() -> None: + bc = VideoBroadcaster(lambda: None, _make_config()) + with pytest.raises(RuntimeError): + await bc.start()