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
77#[non_exhaustive]
79pub enum MediaRole {
80 Ingest(Box<dyn MediaSink>),
82 Egress(Box<EgressSource>),
84}
85
86pub struct Session {
93 rtc: Rtc,
94 socket: Arc<UdpSocket>,
97 locals: Vec<SocketAddr>,
106 inbound: mpsc::Receiver<Packet>,
110 role: MediaRole,
111 writes_rx: Option<mpsc::Receiver<WriteRequest>>,
115 ingest_clock: IngestClock,
118 egress_clock: EgressClock,
121}
122
123impl Session {
124 pub fn ingest(
127 rtc: Rtc,
128 socket: Arc<UdpSocket>,
129 locals: Vec<SocketAddr>,
130 inbound: mpsc::Receiver<Packet>,
131 sink: Box<dyn MediaSink>,
132 ) -> Self {
133 Self {
134 rtc,
135 socket,
136 locals,
137 inbound,
138 role: MediaRole::Ingest(sink),
139 writes_rx: None,
140 ingest_clock: IngestClock::default(),
141 egress_clock: EgressClock::default(),
142 }
143 }
144
145 pub fn egress(
148 rtc: Rtc,
149 socket: Arc<UdpSocket>,
150 locals: Vec<SocketAddr>,
151 inbound: mpsc::Receiver<Packet>,
152 mut source: EgressSource,
153 ) -> Self {
154 let writes_rx = source.take_writes();
155 Self {
156 rtc,
157 socket,
158 locals,
159 inbound,
160 role: MediaRole::Egress(Box::new(source)),
161 writes_rx: Some(writes_rx),
162 ingest_clock: IngestClock::default(),
163 egress_clock: EgressClock::default(),
164 }
165 }
166
167 pub async fn run(mut self) -> Result<()> {
168 let started = Instant::now();
169 let mut connected = false;
170 loop {
171 if !self.rtc.is_alive() {
176 return Err(Error::SessionClosed);
177 }
178
179 if !connected && started.elapsed() >= ICE_ESTABLISH_TIMEOUT {
182 return Err(Error::IceTimeout);
183 }
184
185 let timeout = match self.rtc.poll_output().map_err(Error::Rtc)? {
186 Output::Timeout(t) => t,
187 Output::Transmit(t) => {
188 if let Err(err) = self.socket.send_to(&t.contents, t.destination).await {
189 tracing::warn!(%err, dst = %t.destination, "send failed");
190 }
191 continue;
192 }
193 Output::Event(event) => {
194 if let Event::IceConnectionStateChange(state) = &event {
195 connected |= state.is_connected();
196 }
197 self.handle_event(event)?;
198 continue;
199 }
200 };
201
202 let now = Instant::now();
203 let mut duration = timeout.saturating_duration_since(now);
204 if !connected {
207 duration = duration.min(ICE_ESTABLISH_TIMEOUT.saturating_sub(started.elapsed()));
208 }
209 if duration.is_zero() {
210 self.rtc.handle_input(Input::Timeout(now)).map_err(Error::Rtc)?;
211 continue;
212 }
213
214 tokio::select! {
217 biased;
218
219 Some(req) = async {
222 match self.writes_rx.as_mut() {
223 Some(rx) => rx.recv().await,
224 None => std::future::pending::<Option<WriteRequest>>().await,
225 }
226 } => {
227 let now = Instant::now();
228 let wallclock = self.egress_clock.wallclock(req.time, now);
229 crate::egress::dispatch(&mut self.rtc, req, wallclock);
230 }
231
232 packet = self.inbound.recv() => {
233 match packet {
234 Some((data, src)) => {
235 let now = Instant::now();
236 let local = pick_local(&self.locals, src);
239 let recv = Receive::new(str0m::net::Protocol::Udp, src, local, &data)
240 .map_err(Error::RtcInput)?;
241 self.rtc.handle_input(Input::Receive(now, recv)).map_err(Error::Rtc)?;
242 }
243 None => return Err(Error::SessionClosed),
246 }
247 }
248
249 _ = tokio::time::sleep(duration) => {
250 self.rtc
251 .handle_input(Input::Timeout(Instant::now()))
252 .map_err(Error::Rtc)?;
253 }
254 }
255 }
256 }
257
258 fn handle_event(&mut self, event: Event) -> Result<()> {
259 match event {
260 Event::IceConnectionStateChange(state) => {
261 tracing::debug!(?state, "ice state");
262 if state == IceConnectionState::Disconnected {
263 return Err(Error::SessionClosed);
264 }
265 }
266 Event::MediaAdded(added) => self.handle_media_added(added)?,
267 Event::MediaData(data) => {
268 if let MediaRole::Ingest(sink) = &mut self.role {
272 let media_us = media_time_to_micros(&data.time);
273 let timestamp_us = self.ingest_clock.normalize(data.mid, data.network_time, media_us);
274 sink.on_frame(
275 data.mid,
276 codec::Frame {
277 timestamp_us,
278 payload: bytes::Bytes::from_owner(data.data),
279 },
280 )?;
281 }
282 }
283 Event::SenderFeedback(feedback) => {
284 if matches!(&self.role, MediaRole::Ingest(_)) {
285 self.ingest_clock.observe(feedback.mid, feedback.sender_info);
286 }
287 }
288 Event::KeyframeRequest(req) => {
289 tracing::debug!(?req, "keyframe request from peer");
292 }
293 _ => {}
294 }
295 Ok(())
296 }
297
298 fn handle_media_added(&mut self, added: str0m::media::MediaAdded) -> Result<()> {
299 let pt = self.rtc.media(added.mid).and_then(|m| m.remote_pts().first().copied());
302 let params = pt.and_then(|pt| self.rtc.codec_config().params().iter().find(|p| p.pt() == pt).copied());
303 let params = match params {
304 Some(p) => p,
305 None => {
306 tracing::warn!(?added.mid, "no codec params for media; ignoring");
307 return Ok(());
308 }
309 };
310 let spec = params.spec();
311 let codec = spec.codec;
312
313 match &mut self.role {
314 MediaRole::Ingest(sink) => {
315 let audio_params = if codec.is_audio() {
316 Some((spec.clock_rate.get(), spec.channels.unwrap_or(1) as u32))
317 } else {
318 None
319 };
320 sink.on_track(added.mid, added.kind, codec, audio_params)?;
321 }
322 MediaRole::Egress(source) => {
323 source.on_track(added.mid, codec, params.pt(), spec.clock_rate)?;
324 }
325 }
326 Ok(())
327 }
328}
329
330#[derive(Default)]
344pub(crate) struct IngestClock {
345 arrival_epoch: Option<Instant>,
347 ntp_epoch_us: Option<i128>,
349 tracks: HashMap<str0m::media::Mid, IngestTrackClock>,
350}
351
352impl IngestClock {
353 fn observe(&mut self, mid: str0m::media::Mid, sender: str0m::rtp::rtcp::SenderInfo) {
355 self.tracks.entry(mid).or_default().sender = Some(SenderAnchor::new(sender));
356 self.establish_ntp_epoch();
357 }
358
359 fn normalize(&mut self, mid: str0m::media::Mid, arrival: Instant, media_us: u64) -> u64 {
363 let epoch = *self.arrival_epoch.get_or_insert(arrival);
364 let track = self.tracks.entry(mid).or_default();
365 let offset = *track.arrival_offset_us.get_or_insert_with(|| {
366 let wall_us = if arrival >= epoch {
371 arrival.duration_since(epoch).as_micros() as i64
372 } else {
373 -(epoch.duration_since(arrival).as_micros() as i64)
374 };
375 wall_us as i128 - media_us as i128
376 });
377 let fallback = to_u64(media_us as i128 + offset);
378 let previous = track.last_output_us;
379 track.last_media_us = Some(media_us);
380 track.last_output_us = Some(fallback);
381
382 self.establish_ntp_epoch();
383 let mapped = self
384 .ntp_epoch_us
385 .zip(self.tracks.get(&mid).and_then(|track| track.sender))
386 .map(|(epoch, sender)| to_u64(sender.capture_time_us(media_us) - epoch));
387 let output = match mapped {
388 Some(mapped) => mapped.max(previous.map_or(fallback, |last| last.saturating_add(1))),
392 None => fallback,
393 };
394 self.tracks.get_mut(&mid).expect("track was inserted").last_output_us = Some(output);
395 output
396 }
397
398 fn establish_ntp_epoch(&mut self) {
401 if self.ntp_epoch_us.is_some() || self.tracks.len() < 2 {
405 return;
406 }
407
408 let mut epoch = i128::MAX;
409 for track in self.tracks.values() {
410 let (Some(sender), Some(media_us), Some(output_us)) =
411 (track.sender, track.last_media_us, track.last_output_us)
412 else {
413 return;
414 };
415 epoch = epoch.min(sender.capture_time_us(media_us) - output_us as i128);
419 }
420 self.ntp_epoch_us = Some(epoch);
421 }
422}
423
424#[derive(Default)]
425struct IngestTrackClock {
426 arrival_offset_us: Option<i128>,
427 sender: Option<SenderAnchor>,
428 last_media_us: Option<u64>,
429 last_output_us: Option<u64>,
430}
431
432#[derive(Clone, Copy)]
433struct SenderAnchor {
434 ntp_us: i128,
435 rtp_us: i128,
436}
437
438impl SenderAnchor {
439 fn new(sender: str0m::rtp::rtcp::SenderInfo) -> Self {
440 Self {
441 ntp_us: system_time_to_micros(sender.ntp_time),
442 rtp_us: media_time_to_micros(&sender.rtp_time) as i128,
443 }
444 }
445
446 fn capture_time_us(self, media_us: u64) -> i128 {
447 self.ntp_us + media_us as i128 - self.rtp_us
448 }
449}
450
451fn system_time_to_micros(time: SystemTime) -> i128 {
452 match time.duration_since(UNIX_EPOCH) {
453 Ok(duration) => duration.as_micros() as i128,
454 Err(err) => -(err.duration().as_micros() as i128),
455 }
456}
457
458fn to_u64(value: i128) -> u64 {
459 value.clamp(0, u64::MAX as i128) as u64
460}
461
462pub(crate) fn log_session_end(role: &str, result: &Result<()>) {
467 match result {
468 Ok(()) | Err(Error::SessionClosed) => tracing::debug!(role, "session ended"),
469 Err(Error::IceTimeout) => tracing::debug!(role, "session ended: ICE never connected"),
472 Err(err) => tracing::warn!(%err, role, "session ended"),
473 }
474}
475
476fn pick_local(locals: &[SocketAddr], src: SocketAddr) -> SocketAddr {
481 locals
482 .iter()
483 .find(|l| l.is_ipv4() == src.is_ipv4())
484 .copied()
485 .unwrap_or(locals[0])
486}
487
488fn media_time_to_micros(time: &str0m::media::MediaTime) -> u64 {
490 let numer = time.numer() as i128;
493 let denom = time.denom() as i128;
494 if denom == 0 {
495 return 0;
496 }
497 let micros = (numer.saturating_mul(1_000_000)) / denom;
498 micros.max(0) as u64
499}
500
501pub(crate) struct Bridges {
504 inner: HashMap<str0m::media::Mid, Box<dyn codec::Bridge>>,
505}
506
507impl Bridges {
508 pub fn new() -> Self {
509 Self { inner: HashMap::new() }
510 }
511
512 pub fn insert(&mut self, mid: str0m::media::Mid, bridge: Box<dyn codec::Bridge>) {
513 self.inner.insert(mid, bridge);
514 }
515
516 pub fn push(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()> {
517 if let Some(bridge) = self.inner.get_mut(&mid) {
518 bridge.push(frame)?;
519 }
520 Ok(())
521 }
522}
523
524pub fn rtc_config_with_codecs(codecs: &[str0m::format::Codec]) -> str0m::RtcConfig {
532 use str0m::format::Codec;
533 let mut config = str0m::RtcConfig::new()
539 .clear_codecs()
540 .set_send_buffer_video(EGRESS_SEND_BUFFER_VIDEO);
541 for c in codecs {
542 config = match c {
543 Codec::Opus => config.enable_opus(true),
544 Codec::H264 => config.enable_h264(true),
545 Codec::H265 => config.enable_h265(true),
546 Codec::Vp8 => config.enable_vp8(true),
547 Codec::Vp9 => config.enable_vp9(true),
548 Codec::Av1 => config.enable_av1(true),
549 _ => config,
551 };
552 }
553 config
554}
555
556pub fn rtc_with_codecs(codecs: &[str0m::format::Codec]) -> Rtc {
561 rtc_config_with_codecs(codecs).build(std::time::Instant::now())
562}
563
564pub async fn bind_udp(advertise: &[SocketAddr]) -> Result<(Arc<UdpSocket>, Vec<SocketAddr>)> {
573 let socket = UdpSocket::bind(("0.0.0.0", 0)).await?;
574 let local = socket.local_addr()?;
575 let candidates = advertised_candidates(advertise, local)?;
576 Ok((Arc::new(socket), candidates))
577}
578
579pub(crate) fn advertised_candidates(advertise: &[SocketAddr], local: SocketAddr) -> Result<Vec<SocketAddr>> {
581 let port = local.port();
582 let candidates = if advertise.is_empty() {
583 let ip = match local.ip() {
584 IpAddr::V4(ip) if ip.is_unspecified() => IpAddr::V4(Ipv4Addr::LOCALHOST),
585 IpAddr::V6(ip) if ip.is_unspecified() => IpAddr::V6(Ipv6Addr::LOCALHOST),
586 ip => ip,
587 };
588
589 let candidate = SocketAddr::new(ip, port);
590 if candidate != local {
591 tracing::info!(bound = %local, advertised = %candidate, "webrtc udp bind is unspecified, advertising loopback ICE candidate");
592 }
593 vec![candidate]
594 } else {
595 advertise.iter().map(|addr| SocketAddr::new(addr.ip(), port)).collect()
598 };
599
600 for addr in &candidates {
601 Candidate::host(*addr, "udp").map_err(str0m::RtcError::from)?;
602 }
603 Ok(candidates)
604}
605
606pub fn spawn_socket_reader(socket: Arc<UdpSocket>) -> mpsc::Receiver<Packet> {
610 let (tx, rx) = mpsc::channel(SESSION_INBOX);
611 tokio::spawn(async move {
612 let mut buf = vec![0u8; 65_535];
613 loop {
614 match socket.recv_from(&mut buf).await {
615 Ok((len, src)) => {
618 if let Err(mpsc::error::TrySendError::Closed(_)) = tx.try_send((buf[..len].to_vec(), src)) {
619 break;
620 }
621 }
622 Err(err) => {
623 tracing::warn!(%err, "webrtc client socket recv failed");
624 break;
625 }
626 }
627 }
628 });
629 rx
630}
631
632#[cfg(test)]
633mod tests {
634 use std::time::{Duration, UNIX_EPOCH};
635
636 use str0m::media::Mid;
637 use str0m::rtp::Ssrc;
638 use str0m::rtp::rtcp::SenderInfo;
639
640 use super::*;
641
642 #[test]
643 fn advertised_candidates_use_loopback_for_unspecified_ipv4() {
644 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
645 let candidates = advertised_candidates(&[], local).unwrap();
646 assert_eq!(candidates, vec!["127.0.0.1:4444".parse().unwrap()]);
647 }
648
649 #[test]
650 fn advertised_candidates_use_loopback_for_unspecified_ipv6() {
651 let local: SocketAddr = "[::]:4444".parse().unwrap();
652 let candidates = advertised_candidates(&[], local).unwrap();
653 assert_eq!(candidates, vec!["[::1]:4444".parse().unwrap()]);
654 }
655
656 #[test]
657 fn advertised_candidates_keep_bound_address_when_specific() {
658 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
659 assert_eq!(advertised_candidates(&[], local).unwrap(), vec![local]);
660 }
661
662 #[test]
663 fn advertised_candidates_reuse_bound_port_for_configured_addresses() {
664 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
665 let advertised = vec!["127.0.0.1:1000".parse().unwrap(), "[::1]:2000".parse().unwrap()];
666
667 assert_eq!(
668 advertised_candidates(&advertised, local).unwrap(),
669 vec!["127.0.0.1:4444".parse().unwrap(), "[::1]:4444".parse().unwrap()]
670 );
671 }
672
673 #[test]
674 fn advertised_candidates_reject_configured_unspecified_addresses() {
675 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
676 let advertised = vec!["0.0.0.0:1000".parse().unwrap()];
677 assert!(advertised_candidates(&advertised, local).is_err());
678 }
679
680 #[test]
681 fn pick_local_matches_address_family() {
682 let v4: SocketAddr = "1.2.3.4:5000".parse().unwrap();
683 let v6: SocketAddr = "[2001:db8::1]:5000".parse().unwrap();
684 let locals = vec![v4, v6];
685 let src_v4: SocketAddr = "9.9.9.9:1".parse().unwrap();
686 let src_v6: SocketAddr = "[2001:db8::2]:1".parse().unwrap();
687 assert_eq!(pick_local(&locals, src_v4), v4);
688 assert_eq!(pick_local(&locals, src_v6), v6);
689 assert_eq!(pick_local(&[v4], src_v6), v4);
691 }
692
693 #[test]
694 fn ingest_clock_rebases_first_frame_to_zero() {
695 let mut clock = IngestClock::default();
696 let mid = Mid::from("0");
697 let t0 = Instant::now();
698 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
700 }
701
702 #[test]
703 fn ingest_clock_tracks_rtp_delta_within_track() {
704 let mut clock = IngestClock::default();
705 let mid = Mid::from("0");
706 let t0 = Instant::now();
707 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
708 let arrival = t0 + Duration::from_millis(17); assert_eq!(clock.normalize(mid, arrival, 5_000_020_000), 20_000);
711 }
712
713 #[test]
714 fn ingest_clock_keeps_tracks_in_sync_via_arrival() {
715 let mut clock = IngestClock::default();
716 let audio = Mid::from("0");
717 let video = Mid::from("1");
718 let t0 = Instant::now();
719 assert_eq!(clock.normalize(audio, t0, 1_000_000_000), 0);
721 let video_arrival = t0 + Duration::from_millis(5);
724 assert_eq!(clock.normalize(video, video_arrival, 8_000_000_000), 5_000);
725 assert_eq!(
727 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_000_033_000),
728 38_000
729 );
730 }
731
732 #[test]
733 fn ingest_clock_handles_track_arriving_before_epoch() {
734 let mut clock = IngestClock::default();
735 let audio = Mid::from("0");
736 let video = Mid::from("1");
737 let t0 = Instant::now();
738 assert_eq!(clock.normalize(audio, t0, 1_000_000), 0);
740 let video_arrival = t0 - Duration::from_millis(5);
744 assert_eq!(clock.normalize(video, video_arrival, 8_000_000), 0);
745 assert_eq!(
746 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_033_000),
747 28_000
748 );
749 }
750
751 #[test]
752 fn ingest_clock_replaces_arrival_jitter_with_sender_report_sync() {
753 let mut clock = IngestClock::default();
754 let audio = Mid::from("0");
755 let video = Mid::from("1");
756 let t0 = Instant::now();
757 let audio_base = 1_000_000_000;
758 let video_base = 8_000_000_000;
759
760 assert_eq!(clock.normalize(audio, t0, audio_base), 0);
761 assert_eq!(
762 clock.normalize(video, t0 + Duration::from_millis(50), video_base),
763 50_000
764 );
765 assert_eq!(
766 clock.normalize(audio, t0 + Duration::from_secs(1), audio_base + 1_000_000),
767 1_000_000
768 );
769 assert_eq!(
770 clock.normalize(video, t0 + Duration::from_millis(1_050), video_base + 1_000_000,),
771 1_050_000
772 );
773
774 let report_time = UNIX_EPOCH + Duration::from_secs(1_700_000_001);
777 clock.observe(audio, sender_info(1, report_time, audio_base + 1_000_000));
778 clock.observe(video, sender_info(2, report_time, video_base + 1_000_000));
779
780 let audio_time = clock.normalize(audio, t0 + Duration::from_millis(1_020), audio_base + 1_020_000);
781 let video_time = clock.normalize(video, t0 + Duration::from_millis(1_070), video_base + 1_020_000);
782 assert_eq!(audio_time, video_time);
783 assert_eq!(audio_time, 1_070_000);
784 }
785
786 fn sender_info(ssrc: u32, ntp_time: SystemTime, rtp_us: u64) -> SenderInfo {
787 SenderInfo {
788 ssrc: Ssrc::from(ssrc),
789 ntp_time,
790 rtp_time: str0m::media::MediaTime::from_micros(rtp_us),
791 sender_packet_count: 0,
792 sender_octet_count: 0,
793 }
794 }
795}