Skip to content
Open
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
569 changes: 569 additions & 0 deletions flashdreams/flashdreams/serving/webrtc/encoders.py

Large diffs are not rendered by default.

107 changes: 101 additions & 6 deletions flashdreams/flashdreams/serving/webrtc/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,20 @@
from typing import Any

import torch
from aiortc import RTCConfiguration, RTCPeerConnection, RTCSessionDescription
from aiortc import (
RTCConfiguration,
RTCPeerConnection,
RTCRtpSender,
RTCSessionDescription,
)
from loguru import logger

from flashdreams.serving.webrtc.controls import KeyboardResampler
from flashdreams.serving.webrtc.media import BufferedVideoTrack
from flashdreams.serving.webrtc.encoders import (
DefaultRTCEncoder,
VideoEncoder,
)
from flashdreams.serving.webrtc.media import BufferedVideoTrack, NVENCVideoTrack
from flashdreams.serving.webrtc.server import SessionBusyError
from flashdreams.serving.webrtc.warmup import (
run_loopback_warmup_session,
Expand Down Expand Up @@ -61,7 +70,8 @@ class ManagedWebRTCSession:
"""Per-session state for the single active WebRTC peer connection."""

runtime: Any
video_track: BufferedVideoTrack
video_track: BufferedVideoTrack | NVENCVideoTrack
video_encoder: VideoEncoder
peer_connection: Any
resampler: KeyboardResampler
control_channel: Any | None = None
Expand Down Expand Up @@ -158,6 +168,69 @@ def _make_resampler(self, *, start_v: float) -> KeyboardResampler:
def _register_extra_peer_handlers(self, peer_connection: Any) -> None:
"""Register optional extra peer-connection event handlers."""

def _prefer_h264_video_codec(self, *, transceiver: Any) -> None:
"""Constrain the transceiver's codec preferences to H.264 variants.

Required when the selected encoder emits pre-encoded H.264 packets
(``av.Packet`` route through ``H264Encoder.pack()``): if the SDP
negotiates VP8/VP9 instead, aiortc will pack the H.264 bitstream
under the wrong codec header and the receiver will fail to decode.

If the local aiortc build does not advertise H.264, no preference
is set; the SDP-time fallback in ``_enforce_h264_or_fallback``
will then swap the encoder to :class:`DefaultRTCEncoder`.
"""
caps = RTCRtpSender.getCapabilities("video")
h264_codecs = [c for c in caps.codecs if c.mimeType.lower() == "video/h264"]
if not h264_codecs:
return
transceiver.setCodecPreferences(h264_codecs)

def _enforce_h264_or_fallback(
self,
*,
transceiver: Any,
managed_session: ManagedWebRTCSession,
num_frames: int,
) -> None:
"""Verify H.264 was negotiated; swap to the software encoder if not.

aiortc exposes the negotiated codec set on
``RTCRtpTransceiver._codecs`` after ``setLocalDescription``. We
read it via that attribute (aiortc-internal, but stable in the
pinned version) and, if H.264 did not land, close the hardware
encoder and install a :class:`DefaultRTCEncoder` with a
:class:`BufferedVideoTrack` on the same sender before the first
RTP packet flies. ``replaceTrack`` does not renegotiate; aiortc's
RTP loop will encode raw ``av.VideoFrame`` output with whatever
codec (VP8/VP9/H.264) actually landed in the SDP.
"""
negotiated = getattr(transceiver, "_codecs", None) or []
if negotiated and negotiated[0].mimeType.lower() == "video/h264":
logger.info(
"Video codec negotiated: {} (hardware encoder path active).",
negotiated[0].mimeType,
)
return

chosen = negotiated[0].mimeType if negotiated else "<none>"
logger.warning(
"H.264 preferred by hardware encoder but SDP negotiation "
"landed on {!r}; swapping to the software encoder before "
"streaming begins.",
chosen,
)
# Close the hardware encoder so its NVENC session is released
# promptly; the software adapter has no hardware resources to
# release itself.
managed_session.video_encoder.close()

fallback_encoder = DefaultRTCEncoder(fps=self.fps)
fallback_track = fallback_encoder.create_track(maxsize=num_frames)
transceiver.sender.replaceTrack(fallback_track)
managed_session.video_encoder = fallback_encoder
managed_session.video_track = fallback_track
Comment on lines +226 to +232

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Old NVENCVideoTrack orphaned after encoder fallback

When _enforce_h264_or_fallback swaps the encoder, it overwrites managed_session.video_track with the new BufferedVideoTrack but never calls close() on the old NVENCVideoTrack. ManagedWebRTCSession.close() calls await self.video_track.close() — by that point it reaches the fallback track, not the original NVENC track. The old track's stop() (which sets aiortc's readyState to "ended") is never called, leaving the track live, its asyncio Queue unconsumed, and any internal aiortc observer references pinned until the process exits.


def _on_offer_received(self, offer_sdp: str) -> None:
"""Hook invoked with the remote offer SDP before negotiation."""

Expand Down Expand Up @@ -283,8 +356,17 @@ async def _create_answer_with_runtime_ready_locked(
# frames than steady state; sizing to it would force a per-chunk
# stall, so we size to the steady-state count.
num_frames = self._runtime.peek_steady_chunk_num_frames()
video_track = BufferedVideoTrack(fps=self.fps, maxsize=num_frames)
peer_connection.addTrack(video_track)
video_encoder: VideoEncoder = self._runtime.video_encoder
video_track = video_encoder.create_track(maxsize=num_frames)
# Use ``addTransceiver`` (not ``addTrack``) so we can constrain the
# SDP m-line's codec list via ``setCodecPreferences`` when the
# encoder emits pre-encoded H.264 packets.
video_transceiver = peer_connection.addTransceiver(
video_track,
direction="sendonly",
)
if video_encoder.prefers_codec == "h264":
self._prefer_h264_video_codec(transceiver=video_transceiver)
# Start the resampler's virtual clock at 0; the real anchor is set
# in the ``on_datachannel`` handler so chunk 0's window starts when
# input can actually arrive.
Expand All @@ -293,6 +375,7 @@ async def _create_answer_with_runtime_ready_locked(
managed_session = ManagedWebRTCSession(
runtime=self._runtime,
video_track=video_track,
video_encoder=video_encoder,
peer_connection=peer_connection,
resampler=resampler,
last_client_message_at=loop.time(),
Expand Down Expand Up @@ -350,6 +433,12 @@ async def on_connectionstatechange() -> None:
answer = await peer_connection.createAnswer()
await peer_connection.setLocalDescription(answer)
await wait_for_ice_gathering_complete(peer_connection)
if video_encoder.prefers_codec == "h264":
self._enforce_h264_or_fallback(
transceiver=video_transceiver,
managed_session=managed_session,
num_frames=num_frames,
)
local_description = peer_connection.localDescription
if local_description is None:
raise RuntimeError("Peer connection did not produce local description.")
Expand Down Expand Up @@ -548,6 +637,7 @@ async def _generation_worker(
runtime = managed_session.runtime
resampler = managed_session.resampler
video_track = managed_session.video_track
video_encoder = managed_session.video_encoder

# Stay idle until the user interacts. Generating eagerly would burn
# GPU cycles on a still scene the viewer never sees. Once an event
Expand Down Expand Up @@ -617,7 +707,12 @@ async def _generation_worker(
return
continue
t_after_gen = loop.time()
enqueued = await video_track.enqueue_chunk(result.video_chunk)
delivery = await video_encoder.deliver_chunk(
result.video_chunk,
video_track,
force_keyframe=False,
)
enqueued = delivery.num_frames
t_after_enqueue = loop.time()

gen_ms = (t_after_gen - t_before_gen) * 1e3
Expand Down
116 changes: 116 additions & 0 deletions flashdreams/flashdreams/serving/webrtc/media.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from __future__ import annotations

import asyncio
import contextlib
from collections.abc import Callable
from fractions import Fraction

Expand All @@ -12,6 +13,7 @@
from aiortc import MediaStreamTrack
from aiortc.mediastreams import MediaStreamError
from av import VideoFrame
from av.packet import Packet
from loguru import logger

_STALL_THRESHOLD_MS = 1.0
Expand Down Expand Up @@ -149,3 +151,117 @@ async def close(self) -> None:
break
self._frames.put_nowait(None)
self.stop()


class NVENCVideoTrack(MediaStreamTrack):
"""WebRTC video track that delivers pre-encoded H.264 packets.

Paired with ``PyNvHardwareEncoder``: :meth:`recv` returns
:class:`av.Packet` (not :class:`av.VideoFrame`), which aiortc's
``RTCRtpSender`` routes through ``H264Encoder.pack()`` for RTP
fragmentation only. The encoder sets ``pts`` and ``time_base`` on
each packet before enqueueing; this track only paces delivery to
``fps`` and applies a drop-oldest overflow policy so a slow
consumer cannot stall the encode worker.
"""

kind = "video"

def __init__(self, *, fps: int, maxsize: int) -> None:
super().__init__()
if fps <= 0:
raise ValueError("fps must be > 0")
if maxsize <= 0:
raise ValueError("maxsize must be > 0")
self._fps = fps
self._frame_interval_s = 1.0 / fps
self._next_deadline_s: float | None = None
self._maxsize = maxsize
self._packets: asyncio.Queue[Packet | None] = asyncio.Queue(
maxsize=maxsize,
)
self._closed = False
self._dropped_packets = 0

@property
def fps(self) -> int:
return self._fps

@property
def maxsize(self) -> int:
return self._maxsize

@property
def dropped_packets(self) -> int:
return self._dropped_packets

def qsize(self) -> int:
return self._packets.qsize()

def enqueue_encoded_packet_nowait(self, packet: Packet) -> bool:
"""Synchronously enqueue one packet on the loop thread.

Called from the encode worker via ``loop.call_soon_threadsafe`` so
packets become visible to :meth:`recv` as soon as they are
produced, without waiting for the whole chunk to finish encoding.
Drops the oldest queued packet on overflow so real-time streaming
does not stall behind a slow consumer.
"""
if self._closed:
return False
if self._maxsize > 0 and self._packets.full():
with contextlib.suppress(asyncio.QueueEmpty):
self._packets.get_nowait()
self._dropped_packets += 1
logger.debug(
"NVENCVideoTrack overflow: dropped oldest packet "
"(total dropped={})",
self._dropped_packets,
)
self._packets.put_nowait(packet)
return True

async def recv(self) -> Packet:
if self._closed:
raise MediaStreamError

loop = asyncio.get_running_loop()
t_get_start = loop.time()
packet = await self._packets.get()
if packet is None:
raise MediaStreamError
get_wait_ms = (loop.time() - t_get_start) * 1000.0
first_packet = self._next_deadline_s is None
just_stalled = (not first_packet) and get_wait_ms > _STALL_THRESHOLD_MS

now_s = loop.time()
if first_packet or just_stalled:
self._next_deadline_s = now_s
else:
proposed = self._next_deadline_s + self._frame_interval_s
wait_s = proposed - now_s
if wait_s > 0:
await asyncio.sleep(wait_s)
self._next_deadline_s = proposed
else:
if -wait_s * 1000.0 > _PACING_LAG_LOG_MS:
logger.debug(
"NVENCVideoTrack pacing lag: deadline {:.1f}ms "
"behind walltime; re-anchoring (queue depth {}).",
-wait_s * 1000.0,
self._packets.qsize(),
)
self._next_deadline_s = now_s
return packet

async def close(self) -> None:
if self._closed:
return
self._closed = True
while True:
try:
self._packets.get_nowait()
except asyncio.QueueEmpty:
break
self._packets.put_nowait(None)
self.stop()
Loading
Loading