use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::{Arc, Mutex};
use str0m::ice::{IceCreds, StunMessage};
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use crate::Result;
use crate::session::{Packet, SESSION_INBOX, advertised_candidates};
#[derive(Default)]
struct Registry {
by_ufrag: HashMap<String, mpsc::Sender<Packet>>,
by_addr: HashMap<SocketAddr, mpsc::Sender<Packet>>,
}
pub(crate) struct Mux {
socket: Arc<UdpSocket>,
registry: Arc<Mutex<Registry>>,
candidates: Vec<SocketAddr>,
}
pub(crate) struct Registration {
ufrag: String,
registry: Arc<Mutex<Registry>>,
}
impl Drop for Registration {
fn drop(&mut self) {
let mut registry = self.registry.lock().unwrap();
registry.by_ufrag.remove(&self.ufrag);
registry.by_addr.retain(|_, tx| !tx.is_closed());
}
}
impl Mux {
pub(crate) async fn bind(udp_bind: SocketAddr, ice_candidates: &[SocketAddr]) -> Result<Self> {
let socket = Arc::new(UdpSocket::bind(udp_bind).await?);
let candidates = advertised_candidates(ice_candidates, socket.local_addr()?)?;
let registry = Arc::new(Mutex::new(Registry::default()));
tokio::spawn(demux(socket.clone(), registry.clone()));
tracing::info!(?candidates, bound = %socket.local_addr()?, "webrtc media mux listening");
Ok(Self {
socket,
registry,
candidates,
})
}
pub(crate) fn register(&self) -> (IceCreds, mpsc::Receiver<Packet>, Registration) {
let creds = IceCreds::new();
let (tx, rx) = mpsc::channel(SESSION_INBOX);
self.registry.lock().unwrap().by_ufrag.insert(creds.ufrag.clone(), tx);
let registration = Registration {
ufrag: creds.ufrag.clone(),
registry: self.registry.clone(),
};
(creds, rx, registration)
}
pub(crate) fn socket(&self) -> Arc<UdpSocket> {
self.socket.clone()
}
pub(crate) fn candidates(&self) -> &[SocketAddr] {
&self.candidates
}
}
async fn demux(socket: Arc<UdpSocket>, registry: Arc<Mutex<Registry>>) {
let mut buf = vec![0u8; 65_535];
loop {
let (len, src) = match socket.recv_from(&mut buf).await {
Ok(v) => v,
Err(err) => {
tracing::warn!(%err, "webrtc media mux recv failed");
continue;
}
};
let data = &buf[..len];
let sender = registry.lock().unwrap().by_addr.get(&src).cloned();
let sender = match sender {
Some(sender) => Some(sender),
None => match local_ufrag(data) {
Some(ufrag) => {
let mut registry = registry.lock().unwrap();
match registry.by_ufrag.get(&ufrag).cloned() {
Some(sender) => {
registry.by_addr.insert(src, sender.clone());
Some(sender)
}
None => None,
}
}
None => None,
},
};
let Some(sender) = sender else {
continue;
};
match sender.try_send((data.to_vec(), src)) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {}
Err(mpsc::error::TrySendError::Closed(_)) => {
registry.lock().unwrap().by_addr.remove(&src);
}
}
}
}
fn local_ufrag(data: &[u8]) -> Option<String> {
let msg = StunMessage::parse(data).ok()?;
if !msg.is_binding_request() {
return None;
}
msg.split_username().map(|(local, _remote)| local.to_string())
}