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 let socket_v6 = self.socket.local_addr().map_err(Error::Io)?.is_ipv6();
173 loop {
174 if !self.rtc.is_alive() {
179 return Err(Error::SessionClosed);
180 }
181
182 if !connected && started.elapsed() >= ICE_ESTABLISH_TIMEOUT {
185 return Err(Error::IceTimeout);
186 }
187
188 let timeout = match self.rtc.poll_output().map_err(Error::Rtc)? {
189 Output::Timeout(t) => t,
190 Output::Transmit(t) => {
191 let dst = crate::net::to_family(t.destination, socket_v6);
192 if let Err(err) = self.socket.send_to(&t.contents, dst).await {
193 tracing::warn!(%err, %dst, "send failed");
194 }
195 continue;
196 }
197 Output::Event(event) => {
198 if let Event::IceConnectionStateChange(state) = &event {
199 connected |= state.is_connected();
200 }
201 self.handle_event(event)?;
202 continue;
203 }
204 };
205
206 let now = Instant::now();
207 let mut duration = timeout.saturating_duration_since(now);
208 if !connected {
211 duration = duration.min(ICE_ESTABLISH_TIMEOUT.saturating_sub(started.elapsed()));
212 }
213 if duration.is_zero() {
214 self.rtc.handle_input(Input::Timeout(now)).map_err(Error::Rtc)?;
215 continue;
216 }
217
218 tokio::select! {
221 biased;
222
223 Some(req) = async {
226 match self.writes_rx.as_mut() {
227 Some(rx) => rx.recv().await,
228 None => std::future::pending::<Option<WriteRequest>>().await,
229 }
230 } => {
231 let now = Instant::now();
232 let wallclock = self.egress_clock.wallclock(req.time, now);
233 crate::egress::dispatch(&mut self.rtc, req, wallclock);
234 }
235
236 packet = self.inbound.recv() => {
237 match packet {
238 Some((data, src)) => {
239 let now = Instant::now();
240 let local = pick_local(&self.locals, src);
243 let recv = Receive::new(str0m::net::Protocol::Udp, src, local, &data)
244 .map_err(Error::RtcInput)?;
245 self.rtc.handle_input(Input::Receive(now, recv)).map_err(Error::Rtc)?;
246 }
247 None => return Err(Error::SessionClosed),
250 }
251 }
252
253 _ = tokio::time::sleep(duration) => {
254 self.rtc
255 .handle_input(Input::Timeout(Instant::now()))
256 .map_err(Error::Rtc)?;
257 }
258 }
259 }
260 }
261
262 fn handle_event(&mut self, event: Event) -> Result<()> {
263 match event {
264 Event::IceConnectionStateChange(state) => {
265 tracing::debug!(?state, "ice state");
266 if state == IceConnectionState::Disconnected {
267 return Err(Error::SessionClosed);
268 }
269 }
270 Event::MediaAdded(added) => self.handle_media_added(added)?,
271 Event::MediaData(data) => {
272 if let MediaRole::Ingest(sink) = &mut self.role {
276 let media_us = media_time_to_micros(&data.time);
277 let timestamp_us = self.ingest_clock.normalize(data.mid, data.network_time, media_us);
278 sink.on_frame(
279 data.mid,
280 codec::Frame {
281 timestamp_us,
282 payload: bytes::Bytes::from_owner(data.data),
283 },
284 )?;
285 }
286 }
287 Event::SenderFeedback(feedback) => {
288 if matches!(&self.role, MediaRole::Ingest(_)) {
289 self.ingest_clock.observe(feedback.mid, feedback.sender_info);
290 }
291 }
292 Event::KeyframeRequest(req) => {
293 tracing::debug!(?req, "keyframe request from peer");
296 }
297 _ => {}
298 }
299 Ok(())
300 }
301
302 fn handle_media_added(&mut self, added: str0m::media::MediaAdded) -> Result<()> {
303 let pt = self.rtc.media(added.mid).and_then(|m| m.remote_pts().first().copied());
306 let params = pt.and_then(|pt| self.rtc.codec_config().params().iter().find(|p| p.pt() == pt).copied());
307 let params = match params {
308 Some(p) => p,
309 None => {
310 tracing::warn!(?added.mid, "no codec params for media; ignoring");
311 return Ok(());
312 }
313 };
314 let spec = params.spec();
315 let codec = spec.codec;
316
317 match &mut self.role {
318 MediaRole::Ingest(sink) => {
319 let audio_params = if codec.is_audio() {
320 Some((spec.clock_rate.get(), spec.channels.unwrap_or(1) as u32))
321 } else {
322 None
323 };
324 sink.on_track(added.mid, added.kind, codec, audio_params)?;
325 }
326 MediaRole::Egress(source) => {
327 source.on_track(added.mid, codec, params.pt(), spec.clock_rate)?;
328 }
329 }
330 Ok(())
331 }
332}
333
334#[derive(Default)]
348pub(crate) struct IngestClock {
349 arrival_epoch: Option<Instant>,
351 ntp_epoch_us: Option<i128>,
353 tracks: HashMap<str0m::media::Mid, IngestTrackClock>,
354}
355
356impl IngestClock {
357 fn observe(&mut self, mid: str0m::media::Mid, sender: str0m::rtp::rtcp::SenderInfo) {
359 self.tracks.entry(mid).or_default().sender = Some(SenderAnchor::new(sender));
360 self.establish_ntp_epoch();
361 }
362
363 fn normalize(&mut self, mid: str0m::media::Mid, arrival: Instant, media_us: u64) -> u64 {
367 let epoch = *self.arrival_epoch.get_or_insert(arrival);
368 let track = self.tracks.entry(mid).or_default();
369 let offset = *track.arrival_offset_us.get_or_insert_with(|| {
370 let wall_us = if arrival >= epoch {
375 arrival.duration_since(epoch).as_micros() as i64
376 } else {
377 -(epoch.duration_since(arrival).as_micros() as i64)
378 };
379 wall_us as i128 - media_us as i128
380 });
381 let fallback = to_u64(media_us as i128 + offset);
382 let previous = track.last_output_us;
383 track.last_media_us = Some(media_us);
384 track.last_output_us = Some(fallback);
385
386 self.establish_ntp_epoch();
387 let mapped = self
388 .ntp_epoch_us
389 .zip(self.tracks.get(&mid).and_then(|track| track.sender))
390 .map(|(epoch, sender)| to_u64(sender.capture_time_us(media_us) - epoch));
391 let output = match mapped {
392 Some(mapped) => mapped.max(previous.map_or(fallback, |last| last.saturating_add(1))),
396 None => fallback,
397 };
398 self.tracks.get_mut(&mid).expect("track was inserted").last_output_us = Some(output);
399 output
400 }
401
402 fn establish_ntp_epoch(&mut self) {
405 if self.ntp_epoch_us.is_some() || self.tracks.len() < 2 {
409 return;
410 }
411
412 let mut epoch = i128::MAX;
413 for track in self.tracks.values() {
414 let (Some(sender), Some(media_us), Some(output_us)) =
415 (track.sender, track.last_media_us, track.last_output_us)
416 else {
417 return;
418 };
419 epoch = epoch.min(sender.capture_time_us(media_us) - output_us as i128);
423 }
424 self.ntp_epoch_us = Some(epoch);
425 }
426}
427
428#[derive(Default)]
429struct IngestTrackClock {
430 arrival_offset_us: Option<i128>,
431 sender: Option<SenderAnchor>,
432 last_media_us: Option<u64>,
433 last_output_us: Option<u64>,
434}
435
436#[derive(Clone, Copy)]
437struct SenderAnchor {
438 ntp_us: i128,
439 rtp_us: i128,
440}
441
442impl SenderAnchor {
443 fn new(sender: str0m::rtp::rtcp::SenderInfo) -> Self {
444 Self {
445 ntp_us: system_time_to_micros(sender.ntp_time),
446 rtp_us: media_time_to_micros(&sender.rtp_time) as i128,
447 }
448 }
449
450 fn capture_time_us(self, media_us: u64) -> i128 {
451 self.ntp_us + media_us as i128 - self.rtp_us
452 }
453}
454
455fn system_time_to_micros(time: SystemTime) -> i128 {
456 match time.duration_since(UNIX_EPOCH) {
457 Ok(duration) => duration.as_micros() as i128,
458 Err(err) => -(err.duration().as_micros() as i128),
459 }
460}
461
462fn to_u64(value: i128) -> u64 {
463 value.clamp(0, u64::MAX as i128) as u64
464}
465
466pub(crate) fn log_session_end(role: &str, result: &Result<()>) {
471 match result {
472 Ok(()) | Err(Error::SessionClosed) => tracing::debug!(role, "session ended"),
473 Err(Error::IceTimeout) => tracing::debug!(role, "session ended: ICE never connected"),
476 Err(err) => tracing::warn!(%err, role, "session ended"),
477 }
478}
479
480fn pick_local(locals: &[SocketAddr], src: SocketAddr) -> SocketAddr {
485 locals
486 .iter()
487 .find(|l| l.is_ipv4() == src.is_ipv4())
488 .copied()
489 .unwrap_or(locals[0])
490}
491
492fn media_time_to_micros(time: &str0m::media::MediaTime) -> u64 {
494 let numer = time.numer() as i128;
497 let denom = time.denom() as i128;
498 if denom == 0 {
499 return 0;
500 }
501 let micros = (numer.saturating_mul(1_000_000)) / denom;
502 micros.max(0) as u64
503}
504
505pub(crate) struct Bridges {
508 inner: HashMap<str0m::media::Mid, Box<dyn codec::Bridge>>,
509}
510
511impl Bridges {
512 pub fn new() -> Self {
513 Self { inner: HashMap::new() }
514 }
515
516 pub fn insert(&mut self, mid: str0m::media::Mid, bridge: Box<dyn codec::Bridge>) {
517 self.inner.insert(mid, bridge);
518 }
519
520 pub fn push(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()> {
521 if let Some(bridge) = self.inner.get_mut(&mid) {
522 bridge.push(frame)?;
523 }
524 Ok(())
525 }
526}
527
528pub fn rtc_config_with_codecs(codecs: &[str0m::format::Codec]) -> str0m::RtcConfig {
536 use str0m::format::Codec;
537 let mut config = str0m::RtcConfig::new()
543 .clear_codecs()
544 .set_send_buffer_video(EGRESS_SEND_BUFFER_VIDEO);
545 for c in codecs {
546 config = match c {
547 Codec::Opus => config.enable_opus(true),
548 Codec::H264 => config.enable_h264(true),
549 Codec::H265 => config.enable_h265(true),
550 Codec::Vp8 => config.enable_vp8(true),
551 Codec::Vp9 => config.enable_vp9(true),
552 Codec::Av1 => config.enable_av1(true),
553 _ => config,
555 };
556 }
557 config
558}
559
560pub fn rtc_with_codecs(codecs: &[str0m::format::Codec]) -> Rtc {
565 rtc_config_with_codecs(codecs).build(std::time::Instant::now())
566}
567
568pub async fn bind_udp(advertise: &[SocketAddr]) -> Result<(Arc<UdpSocket>, Vec<SocketAddr>)> {
577 let socket = UdpSocket::bind(("0.0.0.0", 0)).await?;
578 let local = socket.local_addr()?;
579 let candidates = advertised_candidates(advertise, local)?;
580 Ok((Arc::new(socket), candidates))
581}
582
583pub(crate) fn advertised_candidates(advertise: &[SocketAddr], local: SocketAddr) -> Result<Vec<SocketAddr>> {
585 let port = local.port();
586 let candidates = if advertise.is_empty() {
587 let ip = match local.ip() {
588 IpAddr::V4(ip) if ip.is_unspecified() => IpAddr::V4(Ipv4Addr::LOCALHOST),
589 IpAddr::V6(ip) if ip.is_unspecified() => IpAddr::V6(Ipv6Addr::LOCALHOST),
590 ip => ip,
591 };
592
593 let candidate = SocketAddr::new(ip, port);
594 if candidate != local {
595 tracing::info!(bound = %local, advertised = %candidate, "webrtc udp bind is unspecified, advertising loopback ICE candidate");
596 }
597 vec![candidate]
598 } else {
599 advertise.iter().map(|addr| SocketAddr::new(addr.ip(), port)).collect()
602 };
603
604 for addr in &candidates {
605 Candidate::host(*addr, "udp").map_err(str0m::RtcError::from)?;
606 }
607 Ok(candidates)
608}
609
610pub fn spawn_socket_reader(socket: Arc<UdpSocket>) -> mpsc::Receiver<Packet> {
614 let (tx, rx) = mpsc::channel(SESSION_INBOX);
615 tokio::spawn(async move {
616 let mut buf = vec![0u8; 65_535];
617 loop {
618 match socket.recv_from(&mut buf).await {
619 Ok((len, src)) => {
622 let src = crate::net::canonical(src);
623 if let Err(mpsc::error::TrySendError::Closed(_)) = tx.try_send((buf[..len].to_vec(), src)) {
624 break;
625 }
626 }
627 Err(err) => {
628 tracing::warn!(%err, "webrtc client socket recv failed");
629 break;
630 }
631 }
632 }
633 });
634 rx
635}
636
637#[cfg(test)]
638mod tests {
639 use std::time::{Duration, UNIX_EPOCH};
640
641 use str0m::media::Mid;
642 use str0m::rtp::Ssrc;
643 use str0m::rtp::rtcp::SenderInfo;
644
645 use super::*;
646
647 #[test]
648 fn advertised_candidates_use_loopback_for_unspecified_ipv4() {
649 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
650 let candidates = advertised_candidates(&[], local).unwrap();
651 assert_eq!(candidates, vec!["127.0.0.1:4444".parse().unwrap()]);
652 }
653
654 #[test]
655 fn advertised_candidates_use_loopback_for_unspecified_ipv6() {
656 let local: SocketAddr = "[::]:4444".parse().unwrap();
657 let candidates = advertised_candidates(&[], local).unwrap();
658 assert_eq!(candidates, vec!["[::1]:4444".parse().unwrap()]);
659 }
660
661 #[test]
662 fn advertised_candidates_keep_bound_address_when_specific() {
663 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
664 assert_eq!(advertised_candidates(&[], local).unwrap(), vec![local]);
665 }
666
667 #[test]
668 fn advertised_candidates_reuse_bound_port_for_configured_addresses() {
669 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
670 let advertised = vec!["127.0.0.1:1000".parse().unwrap(), "[::1]:2000".parse().unwrap()];
671
672 assert_eq!(
673 advertised_candidates(&advertised, local).unwrap(),
674 vec!["127.0.0.1:4444".parse().unwrap(), "[::1]:4444".parse().unwrap()]
675 );
676 }
677
678 #[test]
679 fn advertised_candidates_reject_configured_unspecified_addresses() {
680 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
681 let advertised = vec!["0.0.0.0:1000".parse().unwrap()];
682 assert!(advertised_candidates(&advertised, local).is_err());
683 }
684
685 #[test]
686 fn pick_local_matches_address_family() {
687 let v4: SocketAddr = "1.2.3.4:5000".parse().unwrap();
688 let v6: SocketAddr = "[2001:db8::1]:5000".parse().unwrap();
689 let locals = vec![v4, v6];
690 let src_v4: SocketAddr = "9.9.9.9:1".parse().unwrap();
691 let src_v6: SocketAddr = "[2001:db8::2]:1".parse().unwrap();
692 assert_eq!(pick_local(&locals, src_v4), v4);
693 assert_eq!(pick_local(&locals, src_v6), v6);
694 assert_eq!(pick_local(&[v4], src_v6), v4);
696 }
697
698 #[test]
699 fn ingest_clock_rebases_first_frame_to_zero() {
700 let mut clock = IngestClock::default();
701 let mid = Mid::from("0");
702 let t0 = Instant::now();
703 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
705 }
706
707 #[test]
708 fn ingest_clock_tracks_rtp_delta_within_track() {
709 let mut clock = IngestClock::default();
710 let mid = Mid::from("0");
711 let t0 = Instant::now();
712 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
713 let arrival = t0 + Duration::from_millis(17); assert_eq!(clock.normalize(mid, arrival, 5_000_020_000), 20_000);
716 }
717
718 #[test]
719 fn ingest_clock_keeps_tracks_in_sync_via_arrival() {
720 let mut clock = IngestClock::default();
721 let audio = Mid::from("0");
722 let video = Mid::from("1");
723 let t0 = Instant::now();
724 assert_eq!(clock.normalize(audio, t0, 1_000_000_000), 0);
726 let video_arrival = t0 + Duration::from_millis(5);
729 assert_eq!(clock.normalize(video, video_arrival, 8_000_000_000), 5_000);
730 assert_eq!(
732 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_000_033_000),
733 38_000
734 );
735 }
736
737 #[test]
738 fn ingest_clock_handles_track_arriving_before_epoch() {
739 let mut clock = IngestClock::default();
740 let audio = Mid::from("0");
741 let video = Mid::from("1");
742 let t0 = Instant::now();
743 assert_eq!(clock.normalize(audio, t0, 1_000_000), 0);
745 let video_arrival = t0 - Duration::from_millis(5);
749 assert_eq!(clock.normalize(video, video_arrival, 8_000_000), 0);
750 assert_eq!(
751 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_033_000),
752 28_000
753 );
754 }
755
756 #[test]
757 fn ingest_clock_replaces_arrival_jitter_with_sender_report_sync() {
758 let mut clock = IngestClock::default();
759 let audio = Mid::from("0");
760 let video = Mid::from("1");
761 let t0 = Instant::now();
762 let audio_base = 1_000_000_000;
763 let video_base = 8_000_000_000;
764
765 assert_eq!(clock.normalize(audio, t0, audio_base), 0);
766 assert_eq!(
767 clock.normalize(video, t0 + Duration::from_millis(50), video_base),
768 50_000
769 );
770 assert_eq!(
771 clock.normalize(audio, t0 + Duration::from_secs(1), audio_base + 1_000_000),
772 1_000_000
773 );
774 assert_eq!(
775 clock.normalize(video, t0 + Duration::from_millis(1_050), video_base + 1_000_000,),
776 1_050_000
777 );
778
779 let report_time = UNIX_EPOCH + Duration::from_secs(1_700_000_001);
782 clock.observe(audio, sender_info(1, report_time, audio_base + 1_000_000));
783 clock.observe(video, sender_info(2, report_time, video_base + 1_000_000));
784
785 let audio_time = clock.normalize(audio, t0 + Duration::from_millis(1_020), audio_base + 1_020_000);
786 let video_time = clock.normalize(video, t0 + Duration::from_millis(1_070), video_base + 1_020_000);
787 assert_eq!(audio_time, video_time);
788 assert_eq!(audio_time, 1_070_000);
789 }
790
791 fn sender_info(ssrc: u32, ntp_time: SystemTime, rtp_us: u64) -> SenderInfo {
792 SenderInfo {
793 ssrc: Ssrc::from(ssrc),
794 ntp_time,
795 rtp_time: str0m::media::MediaTime::from_micros(rtp_us),
796 sender_packet_count: 0,
797 sender_octet_count: 0,
798 }
799 }
800}