1use std::collections::HashMap;
14use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
15use std::sync::Arc;
16use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
17
18use str0m::{Candidate, Event, IceConnectionState, Input, Output, Rtc, net::Receive};
19use tokio::net::UdpSocket;
20use tokio::sync::mpsc;
21
22use crate::egress::{EgressClock, EgressSource, WriteRequest};
23use crate::{Error, Result, codec};
24
25pub(crate) type Packet = (Vec<u8>, SocketAddr);
30
31pub(crate) const SESSION_INBOX: usize = 256;
35
36const EGRESS_SEND_BUFFER_VIDEO: usize = 3000;
41
42const ICE_ESTABLISH_TIMEOUT: Duration = Duration::from_secs(30);
57
58pub trait MediaSink: Send {
63 fn on_track(
65 &mut self,
66 mid: str0m::media::Mid,
67 kind: str0m::media::MediaKind,
68 codec: str0m::format::Codec,
69 audio_params: Option<(u32, u32)>,
70 ) -> Result<()>;
71
72 fn on_frame(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()>;
75
76 fn abort(&mut self, err: moq_net::Error);
79}
80
81#[non_exhaustive]
83pub enum MediaRole {
84 Ingest(Box<dyn MediaSink>),
86 Egress(Box<EgressSource>),
88}
89
90pub struct Session {
97 rtc: Rtc,
98 socket: Arc<UdpSocket>,
101 locals: Vec<SocketAddr>,
110 inbound: mpsc::Receiver<Packet>,
114 role: MediaRole,
115 writes_rx: Option<mpsc::Receiver<WriteRequest>>,
119 ingest_clock: IngestClock,
122 egress_clock: EgressClock,
125}
126
127impl Session {
128 pub fn ingest(
131 rtc: Rtc,
132 socket: Arc<UdpSocket>,
133 locals: Vec<SocketAddr>,
134 inbound: mpsc::Receiver<Packet>,
135 sink: Box<dyn MediaSink>,
136 ) -> Self {
137 Self {
138 rtc,
139 socket,
140 locals,
141 inbound,
142 role: MediaRole::Ingest(sink),
143 writes_rx: None,
144 ingest_clock: IngestClock::default(),
145 egress_clock: EgressClock::default(),
146 }
147 }
148
149 pub fn egress(
152 rtc: Rtc,
153 socket: Arc<UdpSocket>,
154 locals: Vec<SocketAddr>,
155 inbound: mpsc::Receiver<Packet>,
156 mut source: EgressSource,
157 ) -> Self {
158 let writes_rx = source.take_writes();
159 Self {
160 rtc,
161 socket,
162 locals,
163 inbound,
164 role: MediaRole::Egress(Box::new(source)),
165 writes_rx: Some(writes_rx),
166 ingest_clock: IngestClock::default(),
167 egress_clock: EgressClock::default(),
168 }
169 }
170
171 pub async fn run(mut self) -> Result<()> {
172 let result = self.run_loop().await;
173 if let Err(err) = &result
177 && !matches!(err, Error::SessionClosed | Error::IceTimeout)
178 && let MediaRole::Ingest(sink) = &mut self.role
179 {
180 sink.abort(moq_net::Error::Transport(err.to_string()));
181 }
182 result
183 }
184
185 async fn run_loop(&mut self) -> Result<()> {
186 let started = Instant::now();
187 let mut connected = false;
188 let socket_v6 = self.socket.local_addr().map_err(Error::Io)?.is_ipv6();
191 loop {
192 if !self.rtc.is_alive() {
197 return Err(Error::SessionClosed);
198 }
199
200 if !connected && started.elapsed() >= ICE_ESTABLISH_TIMEOUT {
203 return Err(Error::IceTimeout);
204 }
205
206 let timeout = match self.rtc.poll_output().map_err(Error::Rtc)? {
207 Output::Timeout(t) => t,
208 Output::Transmit(t) => {
209 let dst = crate::net::to_family(t.destination, socket_v6);
210 if let Err(err) = self.socket.send_to(&t.contents, dst).await {
211 tracing::warn!(%err, %dst, "send failed");
212 }
213 continue;
214 }
215 Output::Event(event) => {
216 if let Event::IceConnectionStateChange(state) = &event {
217 connected |= state.is_connected();
218 }
219 self.handle_event(event)?;
220 continue;
221 }
222 };
223
224 let now = Instant::now();
225 let mut duration = timeout.saturating_duration_since(now);
226 if !connected {
229 duration = duration.min(ICE_ESTABLISH_TIMEOUT.saturating_sub(started.elapsed()));
230 }
231 if duration.is_zero() {
232 self.rtc.handle_input(Input::Timeout(now)).map_err(Error::Rtc)?;
233 continue;
234 }
235
236 tokio::select! {
239 biased;
240
241 Some(req) = async {
244 match self.writes_rx.as_mut() {
245 Some(rx) => rx.recv().await,
246 None => std::future::pending::<Option<WriteRequest>>().await,
247 }
248 } => {
249 let now = Instant::now();
250 let wallclock = self.egress_clock.wallclock(req.time, now);
251 crate::egress::dispatch(&mut self.rtc, req, wallclock);
252 }
253
254 packet = self.inbound.recv() => {
255 match packet {
256 Some((data, src)) => {
257 let now = Instant::now();
258 let local = pick_local(&self.locals, src);
261 let recv = Receive::new(str0m::net::Protocol::Udp, src, local, &data)
262 .map_err(Error::RtcInput)?;
263 self.rtc.handle_input(Input::Receive(now, recv)).map_err(Error::Rtc)?;
264 }
265 None => return Err(Error::SessionClosed),
268 }
269 }
270
271 _ = tokio::time::sleep(duration) => {
272 self.rtc
273 .handle_input(Input::Timeout(Instant::now()))
274 .map_err(Error::Rtc)?;
275 }
276 }
277 }
278 }
279
280 fn handle_event(&mut self, event: Event) -> Result<()> {
281 match event {
282 Event::IceConnectionStateChange(state) => {
283 tracing::debug!(?state, "ice state");
284 if state == IceConnectionState::Disconnected {
285 return Err(Error::SessionClosed);
286 }
287 }
288 Event::MediaAdded(added) => self.handle_media_added(added)?,
289 Event::MediaData(data) => {
290 if let MediaRole::Ingest(sink) = &mut self.role {
294 let media_us = media_time_to_micros(&data.time);
295 let timestamp_us = self.ingest_clock.normalize(data.mid, data.network_time, media_us);
296 sink.on_frame(
297 data.mid,
298 codec::Frame {
299 timestamp_us,
300 payload: bytes::Bytes::from_owner(data.data),
301 },
302 )?;
303 }
304 }
305 Event::SenderFeedback(feedback) => {
306 if matches!(&self.role, MediaRole::Ingest(_)) {
307 self.ingest_clock.observe(feedback.mid, feedback.sender_info);
308 }
309 }
310 Event::KeyframeRequest(req) => {
311 tracing::debug!(?req, "keyframe request from peer");
314 }
315 _ => {}
316 }
317 Ok(())
318 }
319
320 fn handle_media_added(&mut self, added: str0m::media::MediaAdded) -> Result<()> {
321 let pt = self.rtc.media(added.mid).and_then(|m| m.remote_pts().first().copied());
324 let params = pt.and_then(|pt| self.rtc.codec_config().params().iter().find(|p| p.pt() == pt).copied());
325 let params = match params {
326 Some(p) => p,
327 None => {
328 tracing::warn!(?added.mid, "no codec params for media; ignoring");
329 return Ok(());
330 }
331 };
332 let spec = params.spec();
333 let codec = spec.codec;
334
335 match &mut self.role {
336 MediaRole::Ingest(sink) => {
337 let audio_params = if codec.is_audio() {
338 Some((spec.clock_rate.get(), spec.channels.unwrap_or(1) as u32))
339 } else {
340 None
341 };
342 sink.on_track(added.mid, added.kind, codec, audio_params)?;
343 }
344 MediaRole::Egress(source) => {
345 source.on_track(added.mid, codec, params.pt(), spec.clock_rate)?;
346 }
347 }
348 Ok(())
349 }
350}
351
352#[derive(Default)]
366pub(crate) struct IngestClock {
367 arrival_epoch: Option<Instant>,
369 ntp_epoch_us: Option<i128>,
371 tracks: HashMap<str0m::media::Mid, IngestTrackClock>,
372}
373
374impl IngestClock {
375 fn observe(&mut self, mid: str0m::media::Mid, sender: str0m::rtp::rtcp::SenderInfo) {
377 self.tracks.entry(mid).or_default().sender = Some(SenderAnchor::new(sender));
378 self.establish_ntp_epoch();
379 }
380
381 fn normalize(&mut self, mid: str0m::media::Mid, arrival: Instant, media_us: u64) -> u64 {
385 let epoch = *self.arrival_epoch.get_or_insert(arrival);
386 let track = self.tracks.entry(mid).or_default();
387 let offset = *track.arrival_offset_us.get_or_insert_with(|| {
388 let wall_us = if arrival >= epoch {
393 arrival.duration_since(epoch).as_micros() as i64
394 } else {
395 -(epoch.duration_since(arrival).as_micros() as i64)
396 };
397 wall_us as i128 - media_us as i128
398 });
399 let fallback = to_u64(media_us as i128 + offset);
400 let previous = track.last_output_us;
401 track.last_media_us = Some(media_us);
402 track.last_output_us = Some(fallback);
403
404 self.establish_ntp_epoch();
405 let mapped = self
406 .ntp_epoch_us
407 .zip(self.tracks.get(&mid).and_then(|track| track.sender))
408 .map(|(epoch, sender)| to_u64(sender.capture_time_us(media_us) - epoch));
409 let output = match mapped {
410 Some(mapped) => mapped.max(previous.map_or(fallback, |last| last.saturating_add(1))),
414 None => fallback,
415 };
416 self.tracks.get_mut(&mid).expect("track was inserted").last_output_us = Some(output);
417 output
418 }
419
420 fn establish_ntp_epoch(&mut self) {
423 if self.ntp_epoch_us.is_some() || self.tracks.len() < 2 {
427 return;
428 }
429
430 let mut epoch = i128::MAX;
431 for track in self.tracks.values() {
432 let (Some(sender), Some(media_us), Some(output_us)) =
433 (track.sender, track.last_media_us, track.last_output_us)
434 else {
435 return;
436 };
437 epoch = epoch.min(sender.capture_time_us(media_us) - output_us as i128);
441 }
442 self.ntp_epoch_us = Some(epoch);
443 }
444}
445
446#[derive(Default)]
447struct IngestTrackClock {
448 arrival_offset_us: Option<i128>,
449 sender: Option<SenderAnchor>,
450 last_media_us: Option<u64>,
451 last_output_us: Option<u64>,
452}
453
454#[derive(Clone, Copy)]
455struct SenderAnchor {
456 ntp_us: i128,
457 rtp_us: i128,
458}
459
460impl SenderAnchor {
461 fn new(sender: str0m::rtp::rtcp::SenderInfo) -> Self {
462 Self {
463 ntp_us: system_time_to_micros(sender.ntp_time),
464 rtp_us: media_time_to_micros(&sender.rtp_time) as i128,
465 }
466 }
467
468 fn capture_time_us(self, media_us: u64) -> i128 {
469 self.ntp_us + media_us as i128 - self.rtp_us
470 }
471}
472
473fn system_time_to_micros(time: SystemTime) -> i128 {
474 match time.duration_since(UNIX_EPOCH) {
475 Ok(duration) => duration.as_micros() as i128,
476 Err(err) => -(err.duration().as_micros() as i128),
477 }
478}
479
480fn to_u64(value: i128) -> u64 {
481 value.clamp(0, u64::MAX as i128) as u64
482}
483
484pub(crate) fn log_session_end(role: &str, result: &Result<()>) {
489 match result {
490 Ok(()) | Err(Error::SessionClosed) => tracing::debug!(role, "session ended"),
491 Err(Error::IceTimeout) => tracing::debug!(role, "session ended: ICE never connected"),
494 Err(err) => tracing::warn!(%err, role, "session ended"),
495 }
496}
497
498fn pick_local(locals: &[SocketAddr], src: SocketAddr) -> SocketAddr {
503 locals
504 .iter()
505 .find(|l| l.is_ipv4() == src.is_ipv4())
506 .copied()
507 .unwrap_or(locals[0])
508}
509
510fn media_time_to_micros(time: &str0m::media::MediaTime) -> u64 {
512 let numer = time.numer() as i128;
515 let denom = time.denom() as i128;
516 if denom == 0 {
517 return 0;
518 }
519 let micros = (numer.saturating_mul(1_000_000)) / denom;
520 micros.max(0) as u64
521}
522
523pub(crate) struct Bridges {
526 inner: HashMap<str0m::media::Mid, Box<dyn codec::Bridge>>,
527}
528
529impl Bridges {
530 pub fn new() -> Self {
531 Self { inner: HashMap::new() }
532 }
533
534 pub fn insert(&mut self, mid: str0m::media::Mid, bridge: Box<dyn codec::Bridge>) {
535 self.inner.insert(mid, bridge);
536 }
537
538 pub fn push(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()> {
539 if let Some(bridge) = self.inner.get_mut(&mid) {
540 bridge.push(frame)?;
541 }
542 Ok(())
543 }
544
545 pub fn abort(&mut self, err: moq_net::Error) {
550 for bridge in std::mem::take(&mut self.inner).into_values() {
551 bridge.abort(err.clone());
552 }
553 }
554}
555
556pub fn rtc_config_with_codecs(codecs: &[str0m::format::Codec]) -> str0m::RtcConfig {
564 use str0m::format::Codec;
565 let mut config = str0m::RtcConfig::new()
571 .clear_codecs()
572 .set_send_buffer_video(EGRESS_SEND_BUFFER_VIDEO);
573 for c in codecs {
574 config = match c {
575 Codec::Opus => config.enable_opus(true),
576 Codec::H264 => config.enable_h264(true),
577 Codec::H265 => config.enable_h265(true),
578 Codec::Vp8 => config.enable_vp8(true),
579 Codec::Vp9 => config.enable_vp9(true),
580 Codec::Av1 => config.enable_av1(true),
581 _ => config,
583 };
584 }
585 config
586}
587
588pub fn rtc_with_codecs(codecs: &[str0m::format::Codec]) -> Rtc {
593 rtc_config_with_codecs(codecs).build(std::time::Instant::now())
594}
595
596pub async fn bind_udp(advertise: &[SocketAddr]) -> Result<(Arc<UdpSocket>, Vec<SocketAddr>)> {
605 let socket = UdpSocket::bind(("0.0.0.0", 0)).await?;
606 let local = socket.local_addr()?;
607 let candidates = advertised_candidates(advertise, local)?;
608 Ok((Arc::new(socket), candidates))
609}
610
611pub(crate) fn advertised_candidates(advertise: &[SocketAddr], local: SocketAddr) -> Result<Vec<SocketAddr>> {
613 let port = local.port();
614 let candidates = if advertise.is_empty() {
615 let ip = match local.ip() {
616 IpAddr::V4(ip) if ip.is_unspecified() => IpAddr::V4(Ipv4Addr::LOCALHOST),
617 IpAddr::V6(ip) if ip.is_unspecified() => IpAddr::V6(Ipv6Addr::LOCALHOST),
618 ip => ip,
619 };
620
621 let candidate = SocketAddr::new(ip, port);
622 if candidate != local {
623 tracing::info!(bound = %local, advertised = %candidate, "webrtc udp bind is unspecified, advertising loopback ICE candidate");
624 }
625 vec![candidate]
626 } else {
627 advertise.iter().map(|addr| SocketAddr::new(addr.ip(), port)).collect()
630 };
631
632 for addr in &candidates {
633 Candidate::host(*addr, "udp").map_err(str0m::RtcError::from)?;
634 }
635 Ok(candidates)
636}
637
638pub fn spawn_socket_reader(socket: Arc<UdpSocket>) -> mpsc::Receiver<Packet> {
642 let (tx, rx) = mpsc::channel(SESSION_INBOX);
643 tokio::spawn(async move {
644 let mut buf = vec![0u8; 65_535];
645 loop {
646 match socket.recv_from(&mut buf).await {
647 Ok((len, src)) => {
650 let src = crate::net::canonical(src);
651 if let Err(mpsc::error::TrySendError::Closed(_)) = tx.try_send((buf[..len].to_vec(), src)) {
652 break;
653 }
654 }
655 Err(err) => {
656 tracing::warn!(%err, "webrtc client socket recv failed");
657 break;
658 }
659 }
660 }
661 });
662 rx
663}
664
665#[cfg(test)]
666mod tests {
667 use std::time::{Duration, UNIX_EPOCH};
668
669 use str0m::media::Mid;
670 use str0m::rtp::Ssrc;
671 use str0m::rtp::rtcp::SenderInfo;
672
673 use super::*;
674
675 #[test]
676 fn advertised_candidates_use_loopback_for_unspecified_ipv4() {
677 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
678 let candidates = advertised_candidates(&[], local).unwrap();
679 assert_eq!(candidates, vec!["127.0.0.1:4444".parse().unwrap()]);
680 }
681
682 #[test]
683 fn advertised_candidates_use_loopback_for_unspecified_ipv6() {
684 let local: SocketAddr = "[::]:4444".parse().unwrap();
685 let candidates = advertised_candidates(&[], local).unwrap();
686 assert_eq!(candidates, vec!["[::1]:4444".parse().unwrap()]);
687 }
688
689 #[test]
690 fn advertised_candidates_keep_bound_address_when_specific() {
691 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
692 assert_eq!(advertised_candidates(&[], local).unwrap(), vec![local]);
693 }
694
695 #[test]
696 fn advertised_candidates_reuse_bound_port_for_configured_addresses() {
697 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
698 let advertised = vec!["127.0.0.1:1000".parse().unwrap(), "[::1]:2000".parse().unwrap()];
699
700 assert_eq!(
701 advertised_candidates(&advertised, local).unwrap(),
702 vec!["127.0.0.1:4444".parse().unwrap(), "[::1]:4444".parse().unwrap()]
703 );
704 }
705
706 #[test]
707 fn advertised_candidates_reject_configured_unspecified_addresses() {
708 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
709 let advertised = vec!["0.0.0.0:1000".parse().unwrap()];
710 assert!(advertised_candidates(&advertised, local).is_err());
711 }
712
713 #[test]
714 fn pick_local_matches_address_family() {
715 let v4: SocketAddr = "1.2.3.4:5000".parse().unwrap();
716 let v6: SocketAddr = "[2001:db8::1]:5000".parse().unwrap();
717 let locals = vec![v4, v6];
718 let src_v4: SocketAddr = "9.9.9.9:1".parse().unwrap();
719 let src_v6: SocketAddr = "[2001:db8::2]:1".parse().unwrap();
720 assert_eq!(pick_local(&locals, src_v4), v4);
721 assert_eq!(pick_local(&locals, src_v6), v6);
722 assert_eq!(pick_local(&[v4], src_v6), v4);
724 }
725
726 #[test]
727 fn ingest_clock_rebases_first_frame_to_zero() {
728 let mut clock = IngestClock::default();
729 let mid = Mid::from("0");
730 let t0 = Instant::now();
731 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
733 }
734
735 #[test]
736 fn ingest_clock_tracks_rtp_delta_within_track() {
737 let mut clock = IngestClock::default();
738 let mid = Mid::from("0");
739 let t0 = Instant::now();
740 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
741 let arrival = t0 + Duration::from_millis(17); assert_eq!(clock.normalize(mid, arrival, 5_000_020_000), 20_000);
744 }
745
746 #[test]
747 fn ingest_clock_keeps_tracks_in_sync_via_arrival() {
748 let mut clock = IngestClock::default();
749 let audio = Mid::from("0");
750 let video = Mid::from("1");
751 let t0 = Instant::now();
752 assert_eq!(clock.normalize(audio, t0, 1_000_000_000), 0);
754 let video_arrival = t0 + Duration::from_millis(5);
757 assert_eq!(clock.normalize(video, video_arrival, 8_000_000_000), 5_000);
758 assert_eq!(
760 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_000_033_000),
761 38_000
762 );
763 }
764
765 #[test]
766 fn ingest_clock_handles_track_arriving_before_epoch() {
767 let mut clock = IngestClock::default();
768 let audio = Mid::from("0");
769 let video = Mid::from("1");
770 let t0 = Instant::now();
771 assert_eq!(clock.normalize(audio, t0, 1_000_000), 0);
773 let video_arrival = t0 - Duration::from_millis(5);
777 assert_eq!(clock.normalize(video, video_arrival, 8_000_000), 0);
778 assert_eq!(
779 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_033_000),
780 28_000
781 );
782 }
783
784 #[test]
785 fn ingest_clock_replaces_arrival_jitter_with_sender_report_sync() {
786 let mut clock = IngestClock::default();
787 let audio = Mid::from("0");
788 let video = Mid::from("1");
789 let t0 = Instant::now();
790 let audio_base = 1_000_000_000;
791 let video_base = 8_000_000_000;
792
793 assert_eq!(clock.normalize(audio, t0, audio_base), 0);
794 assert_eq!(
795 clock.normalize(video, t0 + Duration::from_millis(50), video_base),
796 50_000
797 );
798 assert_eq!(
799 clock.normalize(audio, t0 + Duration::from_secs(1), audio_base + 1_000_000),
800 1_000_000
801 );
802 assert_eq!(
803 clock.normalize(video, t0 + Duration::from_millis(1_050), video_base + 1_000_000,),
804 1_050_000
805 );
806
807 let report_time = UNIX_EPOCH + Duration::from_secs(1_700_000_001);
810 clock.observe(audio, sender_info(1, report_time, audio_base + 1_000_000));
811 clock.observe(video, sender_info(2, report_time, video_base + 1_000_000));
812
813 let audio_time = clock.normalize(audio, t0 + Duration::from_millis(1_020), audio_base + 1_020_000);
814 let video_time = clock.normalize(video, t0 + Duration::from_millis(1_070), video_base + 1_020_000);
815 assert_eq!(audio_time, video_time);
816 assert_eq!(audio_time, 1_070_000);
817 }
818
819 fn sender_info(ssrc: u32, ntp_time: SystemTime, rtp_us: u64) -> SenderInfo {
820 SenderInfo {
821 ssrc: Ssrc::from(ssrc),
822 ntp_time,
823 rtp_time: str0m::media::MediaTime::from_micros(rtp_us),
824 sender_packet_count: 0,
825 sender_octet_count: 0,
826 }
827 }
828}