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
164 changes: 149 additions & 15 deletions teleoprtc/stream.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,12 @@
H264RtpPacketizer,
IceServer,
NalUnit,
OpusRtpPacketizer,
OpusRtpDepacketizer,
PeerConnection,
PliHandler,
RtcpNackResponder,
RtcpReceivingSession,
RtcpSrReporter,
RtpPacketizationConfig,
Track,
Expand All @@ -43,6 +46,11 @@ class RTCSessionDescription:


class WebRTCBaseStream(abc.ABC):
# destorying wrapper on close can cause deadlock
# TODO: upstream a fix to this
_retained_messaging_channels: List[DataChannel] = []
_retain_messaging_channel_on_close = False

def __init__(self,
consumed_camera_types: List[str],
consume_audio: bool,
Expand Down Expand Up @@ -71,8 +79,10 @@ def __init__(self,
self.messaging_channel: Optional[DataChannel] = None
self.incoming_message_handlers: List[MessageHandler] = []
self._consumer_tracks: List[Track] = []
self._offered_tracks: Dict[str, Track] = {}
self._incoming_audio_handlers: List[Any] = []
self._sender_tasks: List[asyncio.Task] = []
self._track_state: List[Tuple[Track, TiciVideoStreamTrack, RtpPacketizationConfig]] = []
self._track_state: List[Tuple[Track, Any, RtpPacketizationConfig]] = []
self._receiver_reports: Dict[str, RtcpReceiverReport] = {}
self._receiver_report_tracks: Dict[str, Tuple[Track, int]] = {}

Expand Down Expand Up @@ -120,9 +130,17 @@ def _add_consumer_transceivers(self):
media = Description.Audio("audio", Description.Direction.RecvOnly)
media.add_opus_codec(111)
track = self.peer_connection.add_track(media)
self._set_incoming_audio_handlers(track)
self._consumer_tracks.append(track)
self.incoming_audio_tracks.append(track)

def _set_incoming_audio_handlers(self, track: Track) -> None:
depacketizer = OpusRtpDepacketizer()
rtcp_session = RtcpReceivingSession()
depacketizer.add_to_chain(rtcp_session)
track.set_media_handler(depacketizer)
self._incoming_audio_handlers.extend((depacketizer, rtcp_session))

def _find_offer_video(self, remote_sdp: str, used_mids: set[str]) -> Tuple[str, int]:
desc = Description(remote_sdp, Description.Type.Offer)
for i in range(desc.media_count()):
Expand All @@ -136,6 +154,31 @@ def _find_offer_video(self, remote_sdp: str, used_mids: set[str]) -> Tuple[str,
return media.mid(), payload_type
raise ValueError("Remote SDP does not offer H264 video")

def _find_offer_audio(self, remote_sdp: str, used_mids: set[str]) -> Tuple[str, int, Description.Direction]:
desc = Description(remote_sdp, Description.Type.Offer)
for i in range(desc.media_count()):
media = desc.media(i)
if media is None or media.type() != "audio" or media.mid() in used_mids:
continue
if media.direction() not in (Description.Direction.RecvOnly, Description.Direction.SendRecv):
continue
for payload_type in media.payload_types():
with contextlib.suppress(ValueError):
rtp_map = media.rtp_map(payload_type)
if rtp_map is not None and rtp_map.format.upper() == "OPUS":
return media.mid(), payload_type, media.direction()
raise ValueError("Remote SDP does not offer Opus audio")

def _find_track_h264(self, track: Track) -> Tuple[int, str, Optional[str]]:
media = track.description()
for payload_type in media.payload_types():
with contextlib.suppress(ValueError):
rtp_map = media.rtp_map(payload_type)
if rtp_map.format.upper() == "H264":
profile = rtp_map.fmtps[0] if rtp_map.fmtps else None
return payload_type, rtp_map.format, profile
raise ValueError("Track does not offer H264 video")

def _make_video_media(self, track: TiciVideoStreamTrack, remote_sdp: str, used_mids: set[str]) -> Tuple[Description.Video, int, int, str]:
mid, payload_type = self._find_offer_video(remote_sdp, used_mids)
used_mids.add(mid)
Expand All @@ -151,7 +194,13 @@ def _add_producer_tracks(self, remote_sdp: Optional[str] = None):
used_mids: set[str] = set()
for track in self.outgoing_video_tracks:
media, ssrc, payload_type, cname = self._make_video_media(track, remote_sdp or "", used_mids)
rtc_track = self.peer_connection.add_track(media)
rtc_track = self._offered_tracks.pop(media.mid(), None)
if rtc_track is None:
rtc_track = self.peer_connection.add_track(media)
else:
offered_media = rtc_track.description()
offered_media.add_ssrc(ssrc, cname, f"stream-{random.getrandbits(32):08x}", track.id)
rtc_track.set_description(offered_media)

rtp_config = RtpPacketizationConfig(ssrc, cname, payload_type, H264RtpPacketizer.CLOCK_RATE)
rtp_config.start_timestamp = random.randint(0, 0xFFFFFFFF)
Expand All @@ -169,8 +218,61 @@ def _add_producer_tracks(self, remote_sdp: Optional[str] = None):
self._receiver_report_tracks[camera_type] = (rtc_track, ssrc)
self._track_state.append((rtc_track, track, rtp_config))

if self.outgoing_audio_tracks:
raise NotImplementedError("Audio producer tracks are not implemented with libdatachannel")
for track in self.outgoing_audio_tracks:
if remote_sdp is None:
mid, payload_type = "audio", 111
direction = Description.Direction.SendRecv if self.expected_incoming_audio else Description.Direction.SendOnly
else:
mid, payload_type, offered_direction = self._find_offer_audio(remote_sdp, used_mids)
direction = Description.Direction.SendRecv if offered_direction == Description.Direction.SendRecv else Description.Direction.SendOnly
used_mids.add(mid)

ssrc = random.randint(1, 0xFFFFFFFF)
cname = f"teleoprtc-{random.getrandbits(32):08x}"
stream_id = f"stream-{random.getrandbits(32):08x}"
media = Description.Audio(mid, direction)
media.add_opus_codec(payload_type)
media.add_ssrc(ssrc, cname, stream_id, track.id)
rtc_track = self._offered_tracks.pop(mid, None)
if rtc_track is not None:
rtc_track.set_description(media)
elif remote_sdp is None and self.expected_incoming_audio:
rtc_track = self.incoming_audio_tracks[0]
rtc_track.set_description(media)
else:
rtc_track = self.peer_connection.add_track(media)

rtp_config = RtpPacketizationConfig(ssrc, cname, payload_type, OpusRtpPacketizer.DEFAULT_CLOCK_RATE)
rtp_config.start_timestamp = random.randint(0, 0xFFFFFFFF)
rtp_config.timestamp = rtp_config.start_timestamp
rtp_config.sequence_number = random.randint(0, 0xFFFF)

packetizer = OpusRtpPacketizer(rtp_config)
packetizer.add_to_chain(RtcpSrReporter(rtp_config))
packetizer.add_to_chain(RtcpNackResponder())
if direction == Description.Direction.SendRecv:
# MediaHandler runs outgoing chains front-to-back and incoming chains
# back-to-front. Put receive handlers first so incoming RTP is handled
# by sender RTCP handlers while still packetized, then depacketized as
# the final operation before delivery to Track.receive().
depacketizer = OpusRtpDepacketizer()
rtcp_session = RtcpReceivingSession()
depacketizer.add_to_chain(rtcp_session)
depacketizer.add_to_chain(packetizer)
rtc_track.set_media_handler(depacketizer)
self._incoming_audio_handlers.extend((depacketizer, rtcp_session))
else:
rtc_track.set_media_handler(packetizer)
self._track_state.append((rtc_track, track, rtp_config))

for mid, rtc_track in self._offered_tracks.items():
if rtc_track.description().type() == "video":
# libdatachannel creates local tracks for every remote recvonly section.
# Keep compatibility video MIDs valid, but do not advertise empty streams.
payload_type, codec, profile = self._find_track_h264(rtc_track)
media = Description.Video(mid, Description.Direction.Inactive)
media.add_video_codec(payload_type, codec, profile)
rtc_track.set_description(media)

def _add_messaging_channel(self, channel: Optional[DataChannel] = None):
if channel is None:
Expand All @@ -181,13 +283,32 @@ def on_message(message: Union[bytes, str]):
for handler in list(self.incoming_message_handlers):
self._call_soon_threadsafe(handler, message)

def on_open():
self._set_event(self.messaging_channel_ready_event)

def on_closed():
self._set_event(self.connection_stopped_event)

channel.on_message(on_message)
channel.on_open(lambda: self._set_event(self.messaging_channel_ready_event))
channel.on_closed(lambda: self._set_event(self.connection_stopped_event))
channel.on_open(on_open)
channel.on_closed(on_closed)
if channel.is_open():
self._set_event(self.messaging_channel_ready_event)
self._on_after_media()

def _retain_messaging_channel(self) -> None:
if self.messaging_channel is None:
return

# No native callback can be running before a remote description is set.
if not self.messaging_channel_ready_event.is_set() and self.peer_connection.remote_description() is None:
self.messaging_channel = None
return

if self._retain_messaging_channel_on_close:
self._retained_messaging_channels.append(self.messaging_channel)
self.messaging_channel = None

def _on_connectionstatechange(self, state: PeerConnection.State):
self._log_debug("connection state is %s", state)
if state in (PeerConnection.State.Connected, PeerConnection.State.Failed):
Expand All @@ -202,14 +323,19 @@ def _on_gatheringstatechange(self, state: PeerConnection.GatheringState):

def _on_incoming_track(self, track: Track):
self._log_debug("got track: %s", track.mid())
try:
camera_type, _ = parse_video_track_id(track.mid())
except ValueError:
camera_type = track.mid()
if camera_type in self.expected_incoming_camera_types:
self.incoming_camera_tracks[camera_type] = track
elif self.expected_incoming_audio:
# An offer-created track may be the same transceiver used by an outgoing
# producer. Reusing it avoids adding a duplicate track with the same MID.
self._offered_tracks[track.mid()] = track
if track.description().type() == "audio" and self.expected_incoming_audio:
self._set_incoming_audio_handlers(track)
self.incoming_audio_tracks.append(track)
elif track.description().type() == "video":
try:
camera_type, _ = parse_video_track_id(track.mid())
except ValueError:
camera_type = track.mid()
if camera_type in self.expected_incoming_camera_types:
self.incoming_camera_tracks[camera_type] = track
self._on_after_media()

def _on_incoming_datachannel(self, channel: DataChannel):
Expand Down Expand Up @@ -288,7 +414,7 @@ async def _wait_for_gathering_complete(self):
if self.peer_connection.gathering_state() != PeerConnection.GatheringState.Complete:
await self.gathering_complete_event.wait()

async def _send_track_loop(self, rtc_track: Track, producer_track: TiciVideoStreamTrack, rtp_config: RtpPacketizationConfig):
async def _send_track_loop(self, rtc_track: Track, producer_track: Any, rtp_config: RtpPacketizationConfig):
while True:
if not rtc_track.is_open():
await asyncio.sleep(0.01)
Expand All @@ -301,6 +427,9 @@ async def _send_track_loop(self, rtc_track: Track, producer_track: TiciVideoStre
continue

pts = int(packet.pts or 0)
time_base = getattr(packet, "time_base", None)
if time_base is not None:
pts = int(pts * time_base * rtp_config.clock_rate)
timestamp = (rtp_config.start_timestamp + pts) & 0xFFFFFFFF
rtc_track.send_frame(data, FrameInfo(timestamp))
except asyncio.CancelledError:
Expand Down Expand Up @@ -355,11 +484,13 @@ async def stop(self):
with contextlib.suppress(asyncio.CancelledError):
await task
self._sender_tasks.clear()
self._retain_messaging_channel()
self.peer_connection.close()
self.messaging_channel = None
self.incoming_camera_tracks.clear()
self.incoming_audio_tracks.clear()
self._consumer_tracks.clear()
self._offered_tracks.clear()
self._incoming_audio_handlers.clear()
self._track_state.clear()
self._receiver_reports.clear()
self._receiver_report_tracks.clear()
Expand All @@ -379,6 +510,7 @@ async def start(self) -> RTCSessionDescription:
self._add_consumer_transceivers()
if self.should_add_data_channel:
self._add_messaging_channel()
self._add_producer_tracks()

self.peer_connection.set_local_description(Description.Type.Offer)
await self._wait_for_gathering_complete()
Expand All @@ -398,6 +530,8 @@ async def start(self) -> RTCSessionDescription:


class WebRTCAnswerStream(WebRTCBaseStream):
_retain_messaging_channel_on_close = True

def __init__(self, session: RTCSessionDescription, *args, **kwargs):
super().__init__(*args, **kwargs)
self.session = session
Expand Down
Loading
Loading