use std::future::Future;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use anyhow::{Context, Result, anyhow};
use async_trait::async_trait;
use bytes::Bytes;
use log::debug;
use tokio::net::UdpSocket;
use webrtc_data::data_channel::{Config as DcConfig, DataChannel};
use webrtc_data::message::message_channel_open::ChannelType;
use webrtc_dtls::config::Config as DtlsConfig;
use webrtc_dtls::conn::DTLSConn;
use webrtc_dtls::crypto::Certificate;
use webrtc_sctp::association::{Association, Config as SctpConfig};
use webrtc_sctp::chunk::chunk_payload_data::PayloadProtocolIdentifier;
use webrtc_sctp::stream::ReliabilityType;
use webrtc_util_011::Conn as Conn011;
use wacore::runtime::{AbortHandle, Runtime};
use wacore::voip::engine::TxIdSource;
use wacore::voip::relay_parse::WEB_CLIENT_RELAY_PORT;
use wacore::voip::transport::{
RelayDisconnectReason, RelayTransport, RelayTransportEvent, RelayTransportFactory,
};
const DATA_CHANNEL_LABEL: &str = "pre-negotiated";
const SCTP_PORT: u16 = 5000;
pub use wacore::voip::demux::{RelayPacketKind, classify_relay_packet};
struct DtlsToSctpConn(Arc<DTLSConn>);
fn remap(e: webrtc_util_011::Error) -> webrtc_util::Error {
webrtc_util::Error::Other(e.to_string())
}
#[async_trait]
impl webrtc_util::Conn for DtlsToSctpConn {
async fn connect(&self, addr: SocketAddr) -> Result<(), webrtc_util::Error> {
Conn011::connect(&*self.0, addr).await.map_err(remap)
}
async fn recv(&self, buf: &mut [u8]) -> Result<usize, webrtc_util::Error> {
Conn011::recv(&*self.0, buf).await.map_err(remap)
}
async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr), webrtc_util::Error> {
Conn011::recv_from(&*self.0, buf).await.map_err(remap)
}
async fn send(&self, buf: &[u8]) -> Result<usize, webrtc_util::Error> {
Conn011::send(&*self.0, buf).await.map_err(remap)
}
async fn send_to(&self, buf: &[u8], target: SocketAddr) -> Result<usize, webrtc_util::Error> {
Conn011::send_to(&*self.0, buf, target).await.map_err(remap)
}
fn local_addr(&self) -> Result<SocketAddr, webrtc_util::Error> {
Conn011::local_addr(&*self.0).map_err(remap)
}
fn remote_addr(&self) -> Option<SocketAddr> {
Conn011::remote_addr(&*self.0)
}
async fn close(&self) -> Result<(), webrtc_util::Error> {
Conn011::close(&*self.0).await.map_err(remap)
}
fn as_any(&self) -> &(dyn std::any::Any + Send + Sync) {
self
}
}
pub struct RelayMediaChannel {
dc: Arc<DataChannel>,
pump: std::sync::Mutex<Option<AbortHandle>>,
assoc: std::sync::Mutex<Option<Arc<Association>>>,
runtime: Arc<dyn Runtime>,
}
impl RelayMediaChannel {
#[inline]
fn lock_pump(&self) -> std::sync::MutexGuard<'_, Option<AbortHandle>> {
self.pump.lock().unwrap_or_else(|e| e.into_inner())
}
#[inline]
fn lock_assoc(&self) -> std::sync::MutexGuard<'_, Option<Arc<Association>>> {
self.assoc.lock().unwrap_or_else(|e| e.into_inner())
}
}
impl Drop for RelayMediaChannel {
fn drop(&mut self) {
let assoc = self.lock_assoc().take();
if let Some(assoc) = assoc {
self.runtime
.spawn(Box::pin(async move {
let _ = assoc.close().await;
}))
.detach();
}
}
}
fn install_default_crypto_provider() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
});
}
pub async fn connect_relay_media(
relay_addr: SocketAddr,
runtime: Arc<dyn Runtime>,
) -> Result<RelayMediaChannel> {
let bind_addr = if relay_addr.is_ipv6() {
"[::]:0"
} else {
"0.0.0.0:0"
};
let udp = UdpSocket::bind(bind_addr).await.context("bind udp")?;
udp.connect(relay_addr)
.await
.context("connect udp to relay")?;
let udp: Arc<dyn Conn011 + Send + Sync> = Arc::new(udp);
install_default_crypto_provider();
let cert = Certificate::generate_self_signed(vec!["wa-voip".to_owned()])
.map_err(|e| anyhow!("dtls self-signed cert: {e}"))?;
let dtls_config = DtlsConfig {
certificates: vec![cert],
insecure_skip_verify: true,
server_name: "localhost".to_owned(),
..Default::default()
};
let dtls = DTLSConn::new(udp, dtls_config, true, None)
.await
.map_err(|e| anyhow!("dtls handshake: {e}"))?;
let net_conn: Arc<dyn webrtc_util::Conn + Send + Sync> =
Arc::new(DtlsToSctpConn(Arc::new(dtls)));
let assoc = Association::client(SctpConfig {
net_conn,
max_receive_buffer_size: 0,
max_message_size: 0,
mtu: 0,
name: "wa-voip".to_owned(),
remote_port: SCTP_PORT,
local_port: SCTP_PORT,
})
.await
.map_err(|e| anyhow!("sctp client: {e}"))?;
let assoc = Arc::new(assoc);
let dc = open_media_datachannel(&assoc).await?;
Ok(RelayMediaChannel {
dc: Arc::new(dc),
pump: std::sync::Mutex::new(None),
assoc: std::sync::Mutex::new(Some(assoc)),
runtime,
})
}
async fn open_media_datachannel(assoc: &Arc<Association>) -> Result<DataChannel> {
let stream = assoc
.open_stream(0, PayloadProtocolIdentifier::Binary)
.await
.map_err(|e| anyhow!("open sctp media stream: {e}"))?;
stream.set_reliability_params(true, ReliabilityType::Rexmit, 0);
DataChannel::client(
stream,
DcConfig {
channel_type: ChannelType::PartialReliableRexmitUnordered,
reliability_parameter: 0,
negotiated: true,
label: DATA_CHANNEL_LABEL.to_owned(),
..Default::default()
},
)
.await
.map_err(|e| anyhow!("datachannel client: {e}"))
}
#[derive(Default)]
pub struct RandTxIds;
impl TxIdSource for RandTxIds {
fn next_tx_id(&mut self) -> [u8; 12] {
rand::random()
}
}
#[async_trait]
impl RelayTransport for RelayMediaChannel {
async fn send(&self, data: Bytes) -> Result<()> {
self.dc
.write(&data)
.await
.map(|_| ())
.map_err(|e| anyhow!("relay datachannel write: {e}"))
}
async fn disconnect(&self) {
if let Some(h) = self.lock_pump().take() {
h.abort();
}
let _ = self.dc.close().await;
let assoc = self.lock_assoc().take();
if let Some(assoc) = assoc {
let _ = assoc.close().await;
}
}
async fn reconnect(
&self,
endpoint: SocketAddr,
) -> Result<(
Arc<dyn RelayTransport>,
async_channel::Receiver<RelayTransportEvent>,
)> {
RelayMediaChannelFactory::new(endpoint, self.runtime.clone())
.connect()
.await
}
}
const RELAY_EVENT_CAP: usize = 256;
#[cfg(all(test, feature = "voip-mlow"))]
const RELAY_READ_BUF: usize = 1500;
const RELAY_SCTP_READ_BUF: usize = 65536;
const RELAY_CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(12);
pub struct RelayMediaChannelFactory {
addr: SocketAddr,
runtime: Arc<dyn Runtime>,
}
impl RelayMediaChannelFactory {
pub fn new(addr: SocketAddr, runtime: Arc<dyn Runtime>) -> Self {
Self { addr, runtime }
}
}
fn dial_candidates(advertised: SocketAddr) -> Vec<SocketAddr> {
let mut out = vec![advertised];
if advertised.port() != WEB_CLIENT_RELAY_PORT {
out.push(SocketAddr::new(advertised.ip(), WEB_CLIENT_RELAY_PORT));
}
out
}
#[async_trait]
impl RelayTransportFactory for RelayMediaChannelFactory {
async fn connect(
&self,
) -> Result<(
Arc<dyn RelayTransport>,
async_channel::Receiver<RelayTransportEvent>,
)> {
let dialed = dial_candidates(self.addr);
let candidates: Vec<
Pin<Box<dyn Future<Output = Result<(SocketAddr, RelayMediaChannel)>> + Send>>,
> = dialed
.iter()
.map(|&addr| {
let runtime = self.runtime.clone();
Box::pin(async move {
match connect_relay_media(addr, runtime).await {
Ok(chan) => Ok((addr, chan)),
Err(e) => Err(anyhow!("{addr}: {e}")),
}
}) as Pin<Box<dyn Future<Output = _> + Send>>
})
.collect();
let tried = dialed
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
.join(", ");
let (winner, chan) = wacore::runtime::timeout(
&*self.runtime,
RELAY_CONNECT_TIMEOUT,
futures::future::select_ok(candidates),
)
.await
.map_err(|_| {
anyhow!(
"relay connect timed out after {RELAY_CONNECT_TIMEOUT:?} \
(DTLS/SCTP didn't complete); tried {tried}"
)
})?
.map_err(|e| anyhow!("relay connect: {e}; tried {tried}"))?
.0;
debug!(
"voip: relay media connected on {winner} (advertised {}){}",
self.addr,
if winner == self.addr {
""
} else {
" - relay did NOT serve its advertised port"
}
);
let chan = Arc::new(chan);
let (tx, rx) = async_channel::bounded(RELAY_EVENT_CAP);
let _ = tx.try_send(RelayTransportEvent::Connected);
let pump = self
.runtime
.spawn(Box::pin(relay_read_pump(chan.dc.clone(), tx)));
*chan.lock_pump() = Some(pump);
Ok((chan as Arc<dyn RelayTransport>, rx))
}
}
#[async_trait]
trait RelayChannelRead: Send + Sync {
async fn read_message(&self, buf: &mut [u8]) -> Result<usize, String>;
}
#[async_trait]
impl RelayChannelRead for DataChannel {
async fn read_message(&self, buf: &mut [u8]) -> Result<usize, String> {
self.read(buf).await.map_err(|e| e.to_string())
}
}
async fn relay_read_pump<R: RelayChannelRead>(
dc: Arc<R>,
tx: async_channel::Sender<RelayTransportEvent>,
) {
let mut buf = vec![0u8; RELAY_SCTP_READ_BUF];
loop {
match dc.read_message(&mut buf).await {
Ok(0) => {
let _ = tx
.send(RelayTransportEvent::Disconnected(
RelayDisconnectReason::Closed,
))
.await;
break;
}
Ok(n) => {
match tx.try_send(RelayTransportEvent::PacketReceived(Bytes::copy_from_slice(
&buf[..n],
))) {
Ok(()) => {}
Err(async_channel::TrySendError::Full(ev)) => {
if classify_relay_packet(&buf[..n]) == RelayPacketKind::Stun
&& tx.send(ev).await.is_err()
{
break;
}
}
Err(async_channel::TrySendError::Closed(_)) => break,
}
}
Err(e) => {
let _ = tx
.send(RelayTransportEvent::Disconnected(
RelayDisconnectReason::ReadError(e),
))
.await;
break;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::VecDeque;
use std::sync::Mutex;
struct ScriptedReader {
msgs: Mutex<VecDeque<Result<Vec<u8>, String>>>,
}
impl ScriptedReader {
fn new(msgs: impl IntoIterator<Item = Result<Vec<u8>, String>>) -> Arc<Self> {
Arc::new(Self {
msgs: Mutex::new(msgs.into_iter().collect()),
})
}
}
#[async_trait]
impl RelayChannelRead for ScriptedReader {
async fn read_message(&self, buf: &mut [u8]) -> Result<usize, String> {
match self.msgs.lock().unwrap().pop_front() {
Some(Ok(data)) => {
buf[..data.len()].copy_from_slice(&data);
Ok(data.len())
}
Some(Err(e)) => Err(e),
None => Ok(0), }
}
}
#[test]
fn dial_candidates_adds_the_web_client_port() {
let advertised: SocketAddr = "57.144.43.57:3478".parse().unwrap();
assert_eq!(
dial_candidates(advertised),
vec![
advertised,
"57.144.43.57:3480".parse::<SocketAddr>().unwrap()
],
"the advertised port is dialled first, with 3480 raced alongside it"
);
}
#[test]
fn dial_candidates_does_not_duplicate_the_web_client_port() {
let advertised: SocketAddr = "170.78.54.98:3480".parse().unwrap();
assert_eq!(dial_candidates(advertised), vec![advertised]);
}
#[tokio::test]
async fn pump_maps_reads_then_eof_to_disconnect() {
let reader = ScriptedReader::new([Ok(vec![1, 2, 3]), Ok(vec![4, 5])]);
let (tx, rx) = async_channel::unbounded();
relay_read_pump(reader, tx).await;
match rx.try_recv() {
Ok(RelayTransportEvent::PacketReceived(b)) => assert_eq!(b.as_ref(), &[1u8, 2, 3][..]),
other => panic!("expected first packet, got {other:?}"),
}
match rx.try_recv() {
Ok(RelayTransportEvent::PacketReceived(b)) => assert_eq!(b.as_ref(), &[4u8, 5][..]),
other => panic!("expected second packet, got {other:?}"),
}
assert!(matches!(
rx.try_recv(),
Ok(RelayTransportEvent::Disconnected(
RelayDisconnectReason::Closed
))
));
assert!(rx.try_recv().is_err(), "no events after Disconnected");
}
#[tokio::test]
async fn pump_maps_read_error_to_disconnect() {
let reader = ScriptedReader::new([Ok(vec![9]), Err("relay read failed".to_string())]);
let (tx, rx) = async_channel::unbounded();
relay_read_pump(reader, tx).await;
assert!(matches!(
rx.try_recv(),
Ok(RelayTransportEvent::PacketReceived(_))
));
match rx.try_recv() {
Ok(RelayTransportEvent::Disconnected(RelayDisconnectReason::ReadError(e))) => {
assert_eq!(e, "relay read failed");
}
other => panic!("expected Disconnected(ReadError), got {other:?}"),
}
}
#[tokio::test]
async fn pump_stops_when_receiver_closed() {
let reader = ScriptedReader::new([Ok(vec![1]), Ok(vec![2]), Ok(vec![3])]);
let (tx, rx) = async_channel::unbounded();
rx.close(); relay_read_pump(reader, tx).await;
assert!(rx.try_recv().is_err());
}
#[tokio::test]
async fn pump_preserves_stun_but_drops_media_under_backpressure() {
let reader = ScriptedReader::new([
Ok(vec![0x90, 0x78, 1, 2]), Ok(vec![0x90, 0x78, 3, 4]), Ok(vec![0x00, 0x01, 5, 6]), ]);
let (tx, rx) = async_channel::bounded(1);
let pump = tokio::spawn(relay_read_pump(reader, tx));
let first = rx.recv().await.unwrap();
assert!(
matches!(&first, RelayTransportEvent::PacketReceived(d) if d[0] == 0x90),
"first delivered event is the media that filled the channel, got {first:?}"
);
let second = rx.recv().await.unwrap();
assert!(
matches!(&second, RelayTransportEvent::PacketReceived(d)
if classify_relay_packet(d) == RelayPacketKind::Stun),
"the STUN must be preserved while the media behind the first was dropped, got {second:?}"
);
assert!(matches!(
rx.recv().await.unwrap(),
RelayTransportEvent::Disconnected(RelayDisconnectReason::Closed)
));
pump.await.unwrap();
}
}
#[cfg(all(test, feature = "voip-mlow"))]
mod udp_relay_e2e {
use super::*;
use std::time::Duration;
use wacore::voip::engine::{CallConfig, CallEvent, SequentialTxIds};
use wacore::voip::session::{CallDirection, MediaPipeline, MediaPipelineParams};
use wacore::voip::{CallChannels, CallEngine, video_control_channel};
use webrtc_dtls::config::Config as DtlsConfig;
use webrtc_sctp::association::{Association, Config as SctpConfig};
use crate::voip::driver::run_call_tokio;
const SELF_LID: &str = "111111111111111:0@lid";
const PEER_LID: &str = "222222222222222:0@lid";
const SSRC: u32 = 0x5741_0001;
const SAMPLES: u32 = 960;
fn config(relay_addr: SocketAddr) -> CallConfig {
CallConfig {
call_id: "CID".into(),
direction: CallDirection::Incoming,
self_lid: SELF_LID.into(),
peer_lid: PEER_LID.into(),
call_key: (0u8..32).collect(),
ssrc: SSRC,
audio: wacore::voip::AudioConfig::MLOW_PCM,
relay_token: vec![0xAB; 16],
relay_ip: relay_addr.ip().to_string(),
relay_port: relay_addr.port(),
integrity_key: b"relay-key".to_vec(),
warp_mi_tag_len: 4,
enable_media: true,
enable_video: false,
enable_sframe: false,
}
}
async fn accept_relay(udp: UdpSocket) -> Arc<DataChannel> {
let udp: Arc<dyn Conn011 + Send + Sync> = Arc::new(udp);
let cert = Certificate::generate_self_signed(vec!["wa-relay".to_owned()]).unwrap();
let dtls = DTLSConn::new(
udp,
DtlsConfig {
certificates: vec![cert],
insecure_skip_verify: true,
..Default::default()
},
false, None,
)
.await
.expect("relay dtls server handshake");
let net_conn: Arc<dyn webrtc_util::Conn + Send + Sync> =
Arc::new(DtlsToSctpConn(Arc::new(dtls)));
let assoc = Association::server(SctpConfig {
net_conn,
max_receive_buffer_size: 0,
max_message_size: 0,
mtu: 0,
name: "wa-relay".to_owned(),
remote_port: SCTP_PORT,
local_port: SCTP_PORT,
})
.await
.expect("relay sctp server");
let dc = open_media_datachannel(&Arc::new(assoc))
.await
.expect("relay datachannel");
Arc::new(dc)
}
fn peer_rtp(
peer: &mut MediaPipeline,
enc: &mut wacore::voip::mlow::MlowEncoder,
n: u32,
) -> Vec<u8> {
let tone: Vec<f32> = (0..SAMPLES as usize)
.map(|i| 0.3 * ((i as f32 + (n * SAMPLES) as f32) * 0.07).sin())
.collect();
let frame = enc.encode(&tone).expect("mlow encode");
peer.protect_audio(&frame)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn native_transport_relays_packets_over_loopback_udp() {
install_default_crypto_provider();
let body = async {
let server_udp = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let server_addr = server_udp.local_addr().unwrap();
let saw_outbound_rtp = Arc::new(std::sync::atomic::AtomicBool::new(false));
let saw_for_relay = saw_outbound_rtp.clone();
let relay_task = tokio::spawn(async move {
let mut peek = [0u8; RELAY_READ_BUF];
let (_, client_addr) = server_udp.peek_from(&mut peek).await.unwrap();
server_udp.connect(client_addr).await.unwrap();
let dc = accept_relay(server_udp).await;
let call_key: Vec<u8> = (0u8..32).collect();
let mut peer = MediaPipeline::new(&MediaPipelineParams {
call_key: &call_key,
self_lid: PEER_LID,
peer_lid: SELF_LID,
ssrc: SSRC,
samples_per_packet: SAMPLES,
warp_mi_tag_len: 4,
})
.unwrap();
let mut enc = wacore::voip::mlow::MlowEncoder::new();
let mut buf = vec![0u8; RELAY_READ_BUF];
let mut sent_peer_rtp = 0u32;
loop {
let n = match dc.read(&mut buf).await {
Ok(0) | Err(_) => break,
Ok(n) => n,
};
match classify_relay_packet(&buf[..n]) {
RelayPacketKind::Stun
if wacore::voip::stun::stun_message_type(&buf[..n])
== Some(wacore::voip::stun::MSG_ALLOCATE_REQUEST) =>
{
let transaction_id: [u8; 12] =
wacore::voip::stun::stun_transaction_id(&buf[..n])
.expect("allocate transaction id")
.try_into()
.expect("STUN transaction IDs are 12 bytes");
let ok = wacore::voip::stun::encode_stun_request(
wacore::voip::stun::MSG_ALLOCATE_SUCCESS,
&transaction_id,
&[],
None,
false,
);
dc.write(&Bytes::from(ok)).await.unwrap();
}
RelayPacketKind::Rtp => {
saw_for_relay.store(true, std::sync::atomic::Ordering::SeqCst);
while sent_peer_rtp < 2 {
let pkt = peer_rtp(&mut peer, &mut enc, sent_peer_rtp);
dc.write(&Bytes::from(pkt)).await.unwrap();
sent_peer_rtp += 1;
}
}
_ => {}
}
}
let _ = dc.close().await;
});
let factory = RelayMediaChannelFactory::new(
server_addr,
Arc::new(crate::runtime_impl::TokioRuntime),
);
let (transport, relay_events) = factory.connect().await.expect("native relay connect");
let (mic_tx, mic_rx) = async_channel::unbounded::<Vec<i16>>();
let tone: Vec<i16> = (0..SAMPLES as usize)
.map(|i| (8000.0 * (i as f32 * 0.1).sin()) as i16)
.collect();
mic_tx.try_send(tone).unwrap();
let (spk_tx, spk_rx) = async_channel::unbounded::<Vec<i16>>();
let (ev_tx, ev_rx) = async_channel::unbounded::<CallEvent>();
let eng =
CallEngine::new(config(server_addr), Box::new(SequentialTxIds::new())).unwrap();
let driver = tokio::spawn(run_call_tokio(
transport,
relay_events,
CallChannels {
mic: mic_rx,
speaker: spk_tx,
encoded_audio_in: async_channel::bounded(1).1,
encoded_audio_out: async_channel::bounded(1).0,
events: ev_tx,
rekey: None,
video_in: async_channel::bounded(1).1,
video_out: async_channel::bounded(1).0,
video_ctl: video_control_channel().1,
group_ctl: None,
},
eng,
));
let allocated = async {
loop {
if matches!(ev_rx.recv().await, Ok(CallEvent::RelayAllocated)) {
break;
}
}
};
tokio::time::timeout(Duration::from_secs(10), allocated)
.await
.expect("the relay's allocate-success must surface RelayAllocated");
let audible = async {
loop {
if let Ok(frame) = spk_rx.recv().await
&& frame.iter().any(|&s| s != 0)
{
break;
}
}
};
tokio::time::timeout(Duration::from_secs(10), audible)
.await
.expect("peer RTP must decode to audible playout over the real transport");
assert!(
saw_outbound_rtp.load(std::sync::atomic::Ordering::SeqCst),
"the engine's mic frame must reach the relay as outbound RTP"
);
drop(mic_tx);
driver.abort();
if let Ok(joined) = tokio::time::timeout(Duration::from_secs(5), relay_task).await {
joined.expect("relay task panicked after teardown");
}
};
tokio::time::timeout(Duration::from_secs(30), body)
.await
.expect("the loopback relay round-trip must complete within the bound");
}
}