1use std::collections::HashMap;
14use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
15use std::sync::Arc;
16use std::time::{Duration, Instant};
17
18use str0m::{Candidate, Event, IceConnectionState, Input, Output, Rtc, net::Receive};
19use tokio::net::UdpSocket;
20use tokio::sync::mpsc;
21
22use crate::egress::{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 clock: IngestClock,
118}
119
120impl Session {
121 pub fn ingest(
124 rtc: Rtc,
125 socket: Arc<UdpSocket>,
126 locals: Vec<SocketAddr>,
127 inbound: mpsc::Receiver<Packet>,
128 sink: Box<dyn MediaSink>,
129 ) -> Self {
130 Self {
131 rtc,
132 socket,
133 locals,
134 inbound,
135 role: MediaRole::Ingest(sink),
136 writes_rx: None,
137 clock: IngestClock::default(),
138 }
139 }
140
141 pub fn egress(
144 rtc: Rtc,
145 socket: Arc<UdpSocket>,
146 locals: Vec<SocketAddr>,
147 inbound: mpsc::Receiver<Packet>,
148 mut source: EgressSource,
149 ) -> Self {
150 let writes_rx = source.take_writes();
151 Self {
152 rtc,
153 socket,
154 locals,
155 inbound,
156 role: MediaRole::Egress(Box::new(source)),
157 writes_rx: Some(writes_rx),
158 clock: IngestClock::default(),
159 }
160 }
161
162 pub async fn run(mut self) -> Result<()> {
163 let started = Instant::now();
164 let mut connected = false;
165 loop {
166 if !self.rtc.is_alive() {
171 return Err(Error::SessionClosed);
172 }
173
174 if !connected && started.elapsed() >= ICE_ESTABLISH_TIMEOUT {
177 return Err(Error::IceTimeout);
178 }
179
180 let timeout = match self.rtc.poll_output().map_err(Error::Rtc)? {
181 Output::Timeout(t) => t,
182 Output::Transmit(t) => {
183 if let Err(err) = self.socket.send_to(&t.contents, t.destination).await {
184 tracing::warn!(%err, dst = %t.destination, "send failed");
185 }
186 continue;
187 }
188 Output::Event(event) => {
189 if let Event::IceConnectionStateChange(state) = &event {
190 connected |= state.is_connected();
191 }
192 self.handle_event(event)?;
193 continue;
194 }
195 };
196
197 let now = Instant::now();
198 let mut duration = timeout.saturating_duration_since(now);
199 if !connected {
202 duration = duration.min(ICE_ESTABLISH_TIMEOUT.saturating_sub(started.elapsed()));
203 }
204 if duration.is_zero() {
205 self.rtc.handle_input(Input::Timeout(now)).map_err(Error::Rtc)?;
206 continue;
207 }
208
209 tokio::select! {
212 biased;
213
214 Some(req) = async {
217 match self.writes_rx.as_mut() {
218 Some(rx) => rx.recv().await,
219 None => std::future::pending::<Option<WriteRequest>>().await,
220 }
221 } => {
222 crate::egress::dispatch(&mut self.rtc, req, Instant::now());
223 }
224
225 packet = self.inbound.recv() => {
226 match packet {
227 Some((data, src)) => {
228 let now = Instant::now();
229 let local = pick_local(&self.locals, src);
232 let recv = Receive::new(str0m::net::Protocol::Udp, src, local, &data)
233 .map_err(Error::RtcInput)?;
234 self.rtc.handle_input(Input::Receive(now, recv)).map_err(Error::Rtc)?;
235 }
236 None => return Err(Error::SessionClosed),
239 }
240 }
241
242 _ = tokio::time::sleep(duration) => {
243 self.rtc
244 .handle_input(Input::Timeout(Instant::now()))
245 .map_err(Error::Rtc)?;
246 }
247 }
248 }
249 }
250
251 fn handle_event(&mut self, event: Event) -> Result<()> {
252 match event {
253 Event::IceConnectionStateChange(state) => {
254 tracing::debug!(?state, "ice state");
255 if state == IceConnectionState::Disconnected {
256 return Err(Error::SessionClosed);
257 }
258 }
259 Event::MediaAdded(added) => self.handle_media_added(added)?,
260 Event::MediaData(data) => {
261 if let MediaRole::Ingest(sink) = &mut self.role {
265 let media_us = media_time_to_micros(&data.time);
266 let timestamp_us = self.clock.normalize(data.mid, data.network_time, media_us);
267 sink.on_frame(
268 data.mid,
269 codec::Frame {
270 timestamp_us,
271 payload: bytes::Bytes::from_owner(data.data),
272 },
273 )?;
274 }
275 }
276 Event::KeyframeRequest(req) => {
277 tracing::debug!(?req, "keyframe request from peer");
280 }
281 _ => {}
282 }
283 Ok(())
284 }
285
286 fn handle_media_added(&mut self, added: str0m::media::MediaAdded) -> Result<()> {
287 let pt = self.rtc.media(added.mid).and_then(|m| m.remote_pts().first().copied());
290 let params = pt.and_then(|pt| self.rtc.codec_config().params().iter().find(|p| p.pt() == pt).copied());
291 let params = match params {
292 Some(p) => p,
293 None => {
294 tracing::warn!(?added.mid, "no codec params for media; ignoring");
295 return Ok(());
296 }
297 };
298 let spec = params.spec();
299 let codec = spec.codec;
300
301 match &mut self.role {
302 MediaRole::Ingest(sink) => {
303 let audio_params = if codec.is_audio() {
304 Some((spec.clock_rate.get(), spec.channels.unwrap_or(1) as u32))
305 } else {
306 None
307 };
308 sink.on_track(added.mid, added.kind, codec, audio_params)?;
309 }
310 MediaRole::Egress(source) => {
311 source.on_track(added.mid, codec, params.pt(), spec.clock_rate)?;
312 }
313 }
314 Ok(())
315 }
316}
317
318#[derive(Default)]
332pub(crate) struct IngestClock {
333 epoch: Option<Instant>,
335 offsets: HashMap<str0m::media::Mid, i64>,
337}
338
339impl IngestClock {
340 fn normalize(&mut self, mid: str0m::media::Mid, arrival: Instant, media_us: u64) -> u64 {
344 let epoch = *self.epoch.get_or_insert(arrival);
345 let offset = *self.offsets.entry(mid).or_insert_with(|| {
346 let wall_us = if arrival >= epoch {
351 arrival.duration_since(epoch).as_micros() as i64
352 } else {
353 -(epoch.duration_since(arrival).as_micros() as i64)
354 };
355 wall_us - media_us as i64
356 });
357 (media_us as i64 + offset).max(0) as u64
358 }
359}
360
361pub(crate) fn log_session_end(role: &str, result: &Result<()>) {
366 match result {
367 Ok(()) | Err(Error::SessionClosed) => tracing::debug!(role, "session ended"),
368 Err(Error::IceTimeout) => tracing::debug!(role, "session ended: ICE never connected"),
371 Err(err) => tracing::warn!(%err, role, "session ended"),
372 }
373}
374
375fn pick_local(locals: &[SocketAddr], src: SocketAddr) -> SocketAddr {
380 locals
381 .iter()
382 .find(|l| l.is_ipv4() == src.is_ipv4())
383 .copied()
384 .unwrap_or(locals[0])
385}
386
387fn media_time_to_micros(time: &str0m::media::MediaTime) -> u64 {
389 let numer = time.numer() as i128;
392 let denom = time.denom() as i128;
393 if denom == 0 {
394 return 0;
395 }
396 let micros = (numer.saturating_mul(1_000_000)) / denom;
397 micros.max(0) as u64
398}
399
400pub(crate) struct Bridges {
403 inner: HashMap<str0m::media::Mid, Box<dyn codec::Bridge>>,
404}
405
406impl Bridges {
407 pub fn new() -> Self {
408 Self { inner: HashMap::new() }
409 }
410
411 pub fn insert(&mut self, mid: str0m::media::Mid, bridge: Box<dyn codec::Bridge>) {
412 self.inner.insert(mid, bridge);
413 }
414
415 pub fn push(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()> {
416 if let Some(bridge) = self.inner.get_mut(&mid) {
417 bridge.push(frame)?;
418 }
419 Ok(())
420 }
421}
422
423pub fn rtc_config_with_codecs(codecs: &[str0m::format::Codec]) -> str0m::RtcConfig {
431 use str0m::format::Codec;
432 let mut config = str0m::RtcConfig::new()
438 .clear_codecs()
439 .set_send_buffer_video(EGRESS_SEND_BUFFER_VIDEO);
440 for c in codecs {
441 config = match c {
442 Codec::Opus => config.enable_opus(true),
443 Codec::H264 => config.enable_h264(true),
444 Codec::H265 => config.enable_h265(true),
445 Codec::Vp8 => config.enable_vp8(true),
446 Codec::Vp9 => config.enable_vp9(true),
447 Codec::Av1 => config.enable_av1(true),
448 _ => config,
450 };
451 }
452 config
453}
454
455pub fn rtc_with_codecs(codecs: &[str0m::format::Codec]) -> Rtc {
460 rtc_config_with_codecs(codecs).build(std::time::Instant::now())
461}
462
463pub async fn bind_udp(advertise: &[SocketAddr]) -> Result<(Arc<UdpSocket>, Vec<SocketAddr>)> {
472 let socket = UdpSocket::bind(("0.0.0.0", 0)).await?;
473 let local = socket.local_addr()?;
474 let candidates = advertised_candidates(advertise, local)?;
475 Ok((Arc::new(socket), candidates))
476}
477
478pub(crate) fn advertised_candidates(advertise: &[SocketAddr], local: SocketAddr) -> Result<Vec<SocketAddr>> {
480 let port = local.port();
481 let candidates = if advertise.is_empty() {
482 let ip = match local.ip() {
483 IpAddr::V4(ip) if ip.is_unspecified() => IpAddr::V4(Ipv4Addr::LOCALHOST),
484 IpAddr::V6(ip) if ip.is_unspecified() => IpAddr::V6(Ipv6Addr::LOCALHOST),
485 ip => ip,
486 };
487
488 let candidate = SocketAddr::new(ip, port);
489 if candidate != local {
490 tracing::info!(bound = %local, advertised = %candidate, "webrtc udp bind is unspecified, advertising loopback ICE candidate");
491 }
492 vec![candidate]
493 } else {
494 advertise.iter().map(|addr| SocketAddr::new(addr.ip(), port)).collect()
497 };
498
499 for addr in &candidates {
500 Candidate::host(*addr, "udp").map_err(str0m::RtcError::from)?;
501 }
502 Ok(candidates)
503}
504
505pub fn spawn_socket_reader(socket: Arc<UdpSocket>) -> mpsc::Receiver<Packet> {
509 let (tx, rx) = mpsc::channel(SESSION_INBOX);
510 tokio::spawn(async move {
511 let mut buf = vec![0u8; 65_535];
512 loop {
513 match socket.recv_from(&mut buf).await {
514 Ok((len, src)) => {
517 if let Err(mpsc::error::TrySendError::Closed(_)) = tx.try_send((buf[..len].to_vec(), src)) {
518 break;
519 }
520 }
521 Err(err) => {
522 tracing::warn!(%err, "webrtc client socket recv failed");
523 break;
524 }
525 }
526 }
527 });
528 rx
529}
530
531#[cfg(test)]
532mod tests {
533 use std::time::Duration;
534
535 use str0m::media::Mid;
536
537 use super::*;
538
539 #[test]
540 fn advertised_candidates_use_loopback_for_unspecified_ipv4() {
541 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
542 let candidates = advertised_candidates(&[], local).unwrap();
543 assert_eq!(candidates, vec!["127.0.0.1:4444".parse().unwrap()]);
544 }
545
546 #[test]
547 fn advertised_candidates_use_loopback_for_unspecified_ipv6() {
548 let local: SocketAddr = "[::]:4444".parse().unwrap();
549 let candidates = advertised_candidates(&[], local).unwrap();
550 assert_eq!(candidates, vec!["[::1]:4444".parse().unwrap()]);
551 }
552
553 #[test]
554 fn advertised_candidates_keep_bound_address_when_specific() {
555 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
556 assert_eq!(advertised_candidates(&[], local).unwrap(), vec![local]);
557 }
558
559 #[test]
560 fn advertised_candidates_reuse_bound_port_for_configured_addresses() {
561 let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
562 let advertised = vec!["127.0.0.1:1000".parse().unwrap(), "[::1]:2000".parse().unwrap()];
563
564 assert_eq!(
565 advertised_candidates(&advertised, local).unwrap(),
566 vec!["127.0.0.1:4444".parse().unwrap(), "[::1]:4444".parse().unwrap()]
567 );
568 }
569
570 #[test]
571 fn advertised_candidates_reject_configured_unspecified_addresses() {
572 let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
573 let advertised = vec!["0.0.0.0:1000".parse().unwrap()];
574 assert!(advertised_candidates(&advertised, local).is_err());
575 }
576
577 #[test]
578 fn pick_local_matches_address_family() {
579 let v4: SocketAddr = "1.2.3.4:5000".parse().unwrap();
580 let v6: SocketAddr = "[2001:db8::1]:5000".parse().unwrap();
581 let locals = vec![v4, v6];
582 let src_v4: SocketAddr = "9.9.9.9:1".parse().unwrap();
583 let src_v6: SocketAddr = "[2001:db8::2]:1".parse().unwrap();
584 assert_eq!(pick_local(&locals, src_v4), v4);
585 assert_eq!(pick_local(&locals, src_v6), v6);
586 assert_eq!(pick_local(&[v4], src_v6), v4);
588 }
589
590 #[test]
591 fn ingest_clock_rebases_first_frame_to_zero() {
592 let mut clock = IngestClock::default();
593 let mid = Mid::from("0");
594 let t0 = Instant::now();
595 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
597 }
598
599 #[test]
600 fn ingest_clock_tracks_rtp_delta_within_track() {
601 let mut clock = IngestClock::default();
602 let mid = Mid::from("0");
603 let t0 = Instant::now();
604 assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
605 let arrival = t0 + Duration::from_millis(17); assert_eq!(clock.normalize(mid, arrival, 5_000_020_000), 20_000);
608 }
609
610 #[test]
611 fn ingest_clock_keeps_tracks_in_sync_via_arrival() {
612 let mut clock = IngestClock::default();
613 let audio = Mid::from("0");
614 let video = Mid::from("1");
615 let t0 = Instant::now();
616 assert_eq!(clock.normalize(audio, t0, 1_000_000_000), 0);
618 let video_arrival = t0 + Duration::from_millis(5);
621 assert_eq!(clock.normalize(video, video_arrival, 8_000_000_000), 5_000);
622 assert_eq!(
624 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_000_033_000),
625 38_000
626 );
627 }
628
629 #[test]
630 fn ingest_clock_handles_track_arriving_before_epoch() {
631 let mut clock = IngestClock::default();
632 let audio = Mid::from("0");
633 let video = Mid::from("1");
634 let t0 = Instant::now();
635 assert_eq!(clock.normalize(audio, t0, 1_000_000), 0);
637 let video_arrival = t0 - Duration::from_millis(5);
641 assert_eq!(clock.normalize(video, video_arrival, 8_000_000), 0);
642 assert_eq!(
643 clock.normalize(video, video_arrival + Duration::from_millis(33), 8_033_000),
644 28_000
645 );
646 }
647}