Add support to smallWebRTCTransport for receiving screenshare videos

This commit is contained in:
mattie ruth backman
2025-08-14 09:45:46 -04:00
committed by Mattie Ruth
parent 256ecf4d71
commit 40f1f4ff11
2 changed files with 127 additions and 29 deletions

View File

@@ -220,6 +220,7 @@ class SmallWebRTCClient:
self._video_output_track = None self._video_output_track = None
self._audio_input_track: Optional[AudioStreamTrack] = None self._audio_input_track: Optional[AudioStreamTrack] = None
self._video_input_track: Optional[VideoStreamTrack] = None self._video_input_track: Optional[VideoStreamTrack] = None
self._screen_video_track: Optional[VideoStreamTrack] = None
self._params = None self._params = None
self._audio_in_channels = None self._audio_in_channels = None
@@ -273,22 +274,28 @@ class SmallWebRTCClient:
return cv2.cvtColor(frame_array, conversion_code) return cv2.cvtColor(frame_array, conversion_code)
async def read_video_frame(self): async def read_video_frame(self, video_source: str):
"""Read video frames from the WebRTC connection. """Read video frames from the WebRTC connection.
Reads a video frame from the given MediaStreamTrack, converts it to RGB, Reads a video frame from the given MediaStreamTrack, converts it to RGB,
and creates an InputImageRawFrame. and creates an InputImageRawFrame.
Args:
video_source: Video source to capture ("camera" or "screenVideo").
Yields: Yields:
UserImageRawFrame objects containing video data from the peer. UserImageRawFrame objects containing video data from the peer.
""" """
while True: while True:
if self._video_input_track is None: video_track = (
self._video_input_track if video_source == "camera" else self._screen_video_track
)
if video_track is None:
await asyncio.sleep(0.01) await asyncio.sleep(0.01)
continue continue
try: try:
frame = await asyncio.wait_for(self._video_input_track.recv(), timeout=2.0) frame = await asyncio.wait_for(video_track.recv(), timeout=2.0)
except asyncio.TimeoutError: except asyncio.TimeoutError:
if self._webrtc_connection.is_connected(): if self._webrtc_connection.is_connected():
logger.warning("Timeout: No video frame received within the specified time.") logger.warning("Timeout: No video frame received within the specified time.")
@@ -314,6 +321,7 @@ class SmallWebRTCClient:
size=(frame.width, frame.height), size=(frame.width, frame.height),
format="RGB", format="RGB",
) )
image_frame.transport_source = video_source
yield image_frame yield image_frame
@@ -435,6 +443,7 @@ class SmallWebRTCClient:
self._audio_input_track = self._webrtc_connection.audio_input_track() self._audio_input_track = self._webrtc_connection.audio_input_track()
self._video_input_track = self._webrtc_connection.video_input_track() self._video_input_track = self._webrtc_connection.video_input_track()
self._screen_video_track = self._webrtc_connection.screen_video_input_track()
if self._params.audio_out_enabled: if self._params.audio_out_enabled:
self._audio_output_track = RawAudioTrack(sample_rate=self._out_sample_rate) self._audio_output_track = RawAudioTrack(sample_rate=self._out_sample_rate)
self._webrtc_connection.replace_audio_track(self._audio_output_track) self._webrtc_connection.replace_audio_track(self._audio_output_track)
@@ -451,6 +460,7 @@ class SmallWebRTCClient:
"""Handle peer disconnection cleanup.""" """Handle peer disconnection cleanup."""
self._audio_input_track = None self._audio_input_track = None
self._video_input_track = None self._video_input_track = None
self._screen_video_track = None
self._audio_output_track = None self._audio_output_track = None
self._video_output_track = None self._video_output_track = None
@@ -458,6 +468,7 @@ class SmallWebRTCClient:
"""Handle client connection closure.""" """Handle client connection closure."""
self._audio_input_track = None self._audio_input_track = None
self._video_input_track = None self._video_input_track = None
self._screen_video_track = None
self._audio_output_track = None self._audio_output_track = None
self._video_output_track = None self._video_output_track = None
await self._callbacks.on_client_disconnected(self._webrtc_connection) await self._callbacks.on_client_disconnected(self._webrtc_connection)
@@ -514,6 +525,7 @@ class SmallWebRTCInputTransport(BaseInputTransport):
self._params = params self._params = params
self._receive_audio_task = None self._receive_audio_task = None
self._receive_video_task = None self._receive_video_task = None
self._receive_screen_video_task = None
self._image_requests = {} self._image_requests = {}
# Whether we have seen a StartFrame already. # Whether we have seen a StartFrame already.
@@ -546,11 +558,11 @@ class SmallWebRTCInputTransport(BaseInputTransport):
await self._client.setup(self._params, frame) await self._client.setup(self._params, frame)
await self._client.connect() await self._client.connect()
await self.set_transport_ready(frame)
if not self._receive_audio_task and self._params.audio_in_enabled: if not self._receive_audio_task and self._params.audio_in_enabled:
self._receive_audio_task = self.create_task(self._receive_audio()) self._receive_audio_task = self.create_task(self._receive_audio())
if not self._receive_video_task and self._params.video_in_enabled: if not self._receive_video_task and self._params.video_in_enabled:
self._receive_video_task = self.create_task(self._receive_video()) self._receive_video_task = self.create_task(self._receive_video("camera"))
await self.set_transport_ready(frame)
async def _stop_tasks(self): async def _stop_tasks(self):
"""Stop all background tasks.""" """Stop all background tasks."""
@@ -592,10 +604,14 @@ class SmallWebRTCInputTransport(BaseInputTransport):
except Exception as e: except Exception as e:
logger.error(f"{self} exception receiving data: {e.__class__.__name__} ({e})") logger.error(f"{self} exception receiving data: {e.__class__.__name__} ({e})")
async def _receive_video(self): async def _receive_video(self, video_source: str):
"""Background task for receiving video frames from WebRTC.""" """Background task for receiving video frames from WebRTC.
Args:
video_source: Video source to capture ("camera" or "screenVideo").
"""
try: try:
video_iterator = self._client.read_video_frame() video_iterator = self._client.read_video_frame(video_source)
async for video_frame in video_iterator: async for video_frame in video_iterator:
if video_frame: if video_frame:
await self.push_video_frame(video_frame) await self.push_video_frame(video_frame)
@@ -603,18 +619,20 @@ class SmallWebRTCInputTransport(BaseInputTransport):
# Check if there are any pending image requests and create UserImageRawFrame # Check if there are any pending image requests and create UserImageRawFrame
if self._image_requests: if self._image_requests:
for req_id, request_frame in list(self._image_requests.items()): for req_id, request_frame in list(self._image_requests.items()):
# Create UserImageRawFrame using the current video frame if request_frame.video_source == video_source:
image_frame = UserImageRawFrame( # Create UserImageRawFrame using the current video frame
user_id=request_frame.user_id, image_frame = UserImageRawFrame(
request=request_frame, user_id=request_frame.user_id,
image=video_frame.image, request=request_frame,
size=video_frame.size, image=video_frame.image,
format=video_frame.format, size=video_frame.size,
) format=video_frame.format,
# Push the frame to the pipeline )
await self.push_video_frame(image_frame) image_frame.transport_source = video_source
# Remove from pending requests # Push the frame to the pipeline
del self._image_requests[req_id] await self.push_video_frame(image_frame)
# Remove from pending requests
del self._image_requests[req_id]
except Exception as e: except Exception as e:
logger.error(f"{self} exception receiving data: {e.__class__.__name__} ({e})") logger.error(f"{self} exception receiving data: {e.__class__.__name__} ({e})")
@@ -646,9 +664,45 @@ class SmallWebRTCInputTransport(BaseInputTransport):
self._image_requests[request_id] = frame self._image_requests[request_id] = frame
# If we're not already receiving video, try to get a frame now # If we're not already receiving video, try to get a frame now
if not self._receive_video_task and self._params.video_in_enabled: if (
frame.video_source == "camera"
and not self._receive_video_task
and self._params.video_in_enabled
):
# Start video reception if it's not already running # Start video reception if it's not already running
self._receive_video_task = self.create_task(self._receive_video()) self._receive_video_task = self.create_task(self._receive_video("camera"))
elif (
frame.video_source == "screenVideo"
and not self._receive_screen_video_task
and self._params.video_in_enabled
):
# Start screen video reception if it's not already running
self._receive_screen_video_task = self.create_task(self._receive_video("screenVideo"))
async def capture_participant_video(
self,
video_source: str = "camera",
):
"""Capture video from a specific participant.
Args:
video_source: Video source to capture from.
"""
# If we're not already receiving video, try to get a frame now
if (
video_source == "camera"
and not self._receive_video_task
and self._params.video_in_enabled
):
# Start video reception if it's not already running
self._receive_video_task = self.create_task(self._receive_video("camera"))
elif (
video_source == "screenVideo"
and not self._receive_screen_video_task
and self._params.video_in_enabled
):
# Start screen video reception if it's not already running
self._receive_screen_video_task = self.create_task(self._receive_video("screenVideo"))
class SmallWebRTCOutputTransport(BaseOutputTransport): class SmallWebRTCOutputTransport(BaseOutputTransport):
@@ -835,3 +889,18 @@ class SmallWebRTCTransport(BaseTransport):
async def _on_client_disconnected(self, webrtc_connection): async def _on_client_disconnected(self, webrtc_connection):
"""Handle client disconnection events.""" """Handle client disconnection events."""
await self._call_event_handler("on_client_disconnected", webrtc_connection) await self._call_event_handler("on_client_disconnected", webrtc_connection)
async def capture_participant_video(
self,
video_source: str = "camera",
):
"""Capture video from a specific participant.
Args:
participant_id: ID of the participant to capture video from.
framerate: Desired framerate for video capture.
video_source: Video source to capture from.
color_format: Color format for video frames.
"""
if self._input:
await self._input.capture_participant_video(video_source)

View File

@@ -39,6 +39,7 @@ except ModuleNotFoundError as e:
SIGNALLING_TYPE = "signalling" SIGNALLING_TYPE = "signalling"
AUDIO_TRANSCEIVER_INDEX = 0 AUDIO_TRANSCEIVER_INDEX = 0
VIDEO_TRANSCEIVER_INDEX = 1 VIDEO_TRANSCEIVER_INDEX = 1
SCREEN_VIDEO_TRANSCEIVER_INDEX = 2
class TrackStatusMessage(BaseModel): class TrackStatusMessage(BaseModel):
@@ -94,14 +95,16 @@ class SmallWebRTCTrack:
enable/disable control and frame discarding for audio and video streams. enable/disable control and frame discarding for audio and video streams.
""" """
def __init__(self, track: MediaStreamTrack): def __init__(self, track: MediaStreamTrack, index: int):
"""Initialize the WebRTC track wrapper. """Initialize the WebRTC track wrapper.
Args: Args:
track: The underlying MediaStreamTrack to wrap. track: The underlying MediaStreamTrack to wrap.
index: The index of the track in the transceiver (0 for mic, 1 for cam, 2 for screen)
""" """
self._track = track self._track = track
self._enabled = True self._enabled = True
self.source_index = index
def set_enabled(self, enabled: bool) -> None: def set_enabled(self, enabled: bool) -> None:
"""Enable or disable the track. """Enable or disable the track.
@@ -191,6 +194,7 @@ class SmallWebRTCConnection(BaseObject):
self._track_getters = { self._track_getters = {
AUDIO_TRANSCEIVER_INDEX: self.audio_input_track, AUDIO_TRANSCEIVER_INDEX: self.audio_input_track,
VIDEO_TRANSCEIVER_INDEX: self.video_input_track, VIDEO_TRANSCEIVER_INDEX: self.video_input_track,
SCREEN_VIDEO_TRANSCEIVER_INDEX: self.screen_video_input_track,
} }
self._initialize() self._initialize()
@@ -343,6 +347,9 @@ class SmallWebRTCConnection(BaseObject):
video_input_track = self.video_input_track() video_input_track = self.video_input_track()
if video_input_track: if video_input_track:
await self.video_input_track().discard_old_frames() await self.video_input_track().discard_old_frames()
screen_video_input_track = self.screen_video_input_track()
if screen_video_input_track:
await self.screen_video_input_track().discard_old_frames()
self.ask_to_renegotiate() self.ask_to_renegotiate()
async def renegotiate(self, sdp: str, type: str, restart_pc: bool = False): async def renegotiate(self, sdp: str, type: str, restart_pc: bool = False):
@@ -488,15 +495,15 @@ class SmallWebRTCConnection(BaseObject):
return self._track_map[AUDIO_TRANSCEIVER_INDEX] return self._track_map[AUDIO_TRANSCEIVER_INDEX]
# Transceivers always appear in creation-order for both peers # Transceivers always appear in creation-order for both peers
# For now we are only considering that we are going to have 02 transceivers, # For support 3 receivers in the following order:
# one for audio and one for video # audio, video, screenVideo
transceivers = self._pc.getTransceivers() transceivers = self._pc.getTransceivers()
if len(transceivers) == 0 or not transceivers[AUDIO_TRANSCEIVER_INDEX].receiver: if len(transceivers) == 0 or not transceivers[AUDIO_TRANSCEIVER_INDEX].receiver:
logger.warning("No audio transceiver is available") logger.warning("No audio transceiver is available")
return None return None
track = transceivers[AUDIO_TRANSCEIVER_INDEX].receiver.track track = transceivers[AUDIO_TRANSCEIVER_INDEX].receiver.track
audio_track = SmallWebRTCTrack(track) if track else None audio_track = SmallWebRTCTrack(track, AUDIO_TRANSCEIVER_INDEX) if track else None
self._track_map[AUDIO_TRANSCEIVER_INDEX] = audio_track self._track_map[AUDIO_TRANSCEIVER_INDEX] = audio_track
return audio_track return audio_track
@@ -510,18 +517,40 @@ class SmallWebRTCConnection(BaseObject):
return self._track_map[VIDEO_TRANSCEIVER_INDEX] return self._track_map[VIDEO_TRANSCEIVER_INDEX]
# Transceivers always appear in creation-order for both peers # Transceivers always appear in creation-order for both peers
# For now we are only considering that we are going to have 02 transceivers, # For support 3 receivers in the following order:
# one for audio and one for video # audio, video, screenVideo
transceivers = self._pc.getTransceivers() transceivers = self._pc.getTransceivers()
if len(transceivers) <= 1 or not transceivers[VIDEO_TRANSCEIVER_INDEX].receiver: if len(transceivers) <= 1 or not transceivers[VIDEO_TRANSCEIVER_INDEX].receiver:
logger.warning("No video transceiver is available") logger.warning("No video transceiver is available")
return None return None
track = transceivers[VIDEO_TRANSCEIVER_INDEX].receiver.track track = transceivers[VIDEO_TRANSCEIVER_INDEX].receiver.track
video_track = SmallWebRTCTrack(track) if track else None video_track = SmallWebRTCTrack(track, VIDEO_TRANSCEIVER_INDEX) if track else None
self._track_map[VIDEO_TRANSCEIVER_INDEX] = video_track self._track_map[VIDEO_TRANSCEIVER_INDEX] = video_track
return video_track return video_track
def screen_video_input_track(self):
"""Get the screen video input track wrapper.
Returns:
SmallWebRTCTrack wrapper for the screen video track, or None if unavailable.
"""
if self._track_map.get(SCREEN_VIDEO_TRANSCEIVER_INDEX):
return self._track_map[SCREEN_VIDEO_TRANSCEIVER_INDEX]
# Transceivers always appear in creation-order for both peers
# For support 3 receivers in the following order:
# audio, video, screenVideo
transceivers = self._pc.getTransceivers()
if len(transceivers) <= 1 or not transceivers[SCREEN_VIDEO_TRANSCEIVER_INDEX].receiver:
logger.warning("No screen video transceiver is available")
return None
track = transceivers[SCREEN_VIDEO_TRANSCEIVER_INDEX].receiver.track
video_track = SmallWebRTCTrack(track, SCREEN_VIDEO_TRANSCEIVER_INDEX) if track else None
self._track_map[SCREEN_VIDEO_TRANSCEIVER_INDEX] = video_track
return video_track
def send_app_message(self, message: Any): def send_app_message(self, message: Any):
"""Send an application message through the data channel. """Send an application message through the data channel.