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 src = crate::net::canonical(src);
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())
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use std::time::Duration;
use str0m::ice::{StunMessage, TransId};
use super::*;
fn sha1_hmac(key: &[u8], payloads: &[&[u8]]) -> [u8; 20] {
use aws_lc_rs::hmac;
let key = hmac::Key::new(hmac::HMAC_SHA1_FOR_LEGACY_USE_ONLY, key);
let mut ctx = hmac::Context::with_key(&key);
for payload in payloads {
ctx.update(payload);
}
let mut out = [0u8; 20];
out.copy_from_slice(ctx.sign().as_ref());
out
}
fn binding_request(creds: &IceCreds) -> Vec<u8> {
let username = format!("{}:peer", creds.ufrag);
let msg = StunMessage::binding_request(&username, TransId::new(), true, 0, 0, false);
let mut buf = vec![0u8; 1024];
let len = msg
.to_bytes(Some(creds.pass.as_bytes()), &mut buf, sha1_hmac)
.expect("serialize binding request");
buf.truncate(len);
buf
}
#[tokio::test]
async fn dual_stack_mux_routes_an_ipv4_peer_with_a_canonical_source() {
let bind = SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 0);
let advertise = [
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0),
SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 0),
];
let mux = Mux::bind(bind, &advertise).await.expect("dual-stack bind");
let port = mux.socket.local_addr().unwrap().port();
let (creds, mut rx, _registration) = mux.register();
let peer = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let peer_addr = peer.local_addr().unwrap();
peer.send_to(&binding_request(&creds), (Ipv4Addr::LOCALHOST, port))
.await
.expect("IPv4 peer must reach a dual-stack mux");
let (_data, src) = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("the demux must route the IPv4 peer's binding request")
.expect("session inbox stayed open");
assert_eq!(src, peer_addr, "an IPv4 peer's source must arrive unmapped");
assert!(src.is_ipv4());
}
}