use crate::peer_connection::PeerConnectionRef;
use crate::runtime::{Mutex, Receiver};
use bytes::BytesMut;
use futures::FutureExt;
use rtc::interceptor::{Interceptor, NoopInterceptor};
use rtc::shared::error::{Error, Result};
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::time::Duration;
pub use rtc::data_channel::{
RTCDataChannelId, RTCDataChannelInit, RTCDataChannelMessage, RTCDataChannelState,
};
#[async_trait::async_trait]
pub trait DataChannel: Send + Sync + 'static {
async fn label(&self) -> Result<String>;
async fn ordered(&self) -> Result<bool>;
async fn max_packet_life_time(&self) -> Result<Option<u16>>;
async fn max_retransmits(&self) -> Result<Option<u16>>;
async fn protocol(&self) -> Result<String>;
async fn negotiated(&self) -> Result<bool>;
fn id(&self) -> RTCDataChannelId;
async fn ready_state(&self) -> Result<RTCDataChannelState>;
async fn buffered_amount_high_threshold(&self) -> Result<u32>;
async fn set_buffered_amount_high_threshold(&self, threshold: u32) -> Result<()>;
async fn buffered_amount_low_threshold(&self) -> Result<u32>;
async fn set_buffered_amount_low_threshold(&self, threshold: u32) -> Result<()>;
async fn outstanding_bytes(&self) -> Result<usize> {
Ok(0)
}
async fn send(&self, data: BytesMut) -> Result<()>;
async fn send_text(&self, text: &str) -> Result<()>;
async fn writable(&self) -> Result<()> {
Ok(())
}
async fn try_send(&self, data: BytesMut) -> Result<()> {
self.send(data).await
}
async fn try_send_text(&self, text: &str) -> Result<()> {
self.send_text(text).await
}
async fn poll(&self) -> Option<DataChannelEvent>;
async fn close(&self) -> Result<()>;
}
#[derive(Debug, Clone)]
pub enum DataChannelEvent {
OnOpen,
OnError,
OnClosing,
OnClose,
OnBufferedAmountLow,
OnBufferedAmountHigh,
OnMessage(RTCDataChannelMessage),
}
pub(crate) struct DataChannelImpl<I = NoopInterceptor>
where
I: Interceptor,
{
id: RTCDataChannelId,
inner: Arc<PeerConnectionRef<I>>,
evt_rx: Mutex<Receiver<DataChannelEvent>>,
}
impl<I> DataChannelImpl<I>
where
I: Interceptor,
{
pub(crate) fn new(
id: RTCDataChannelId,
inner: Arc<PeerConnectionRef<I>>,
evt_rx: Receiver<DataChannelEvent>,
) -> Self {
Self {
id,
inner,
evt_rx: Mutex::new(evt_rx),
}
}
async fn await_send_capacity(&self) {
futures::select! {
_ = self.inner.data_channel_backpressure.notified().fuse() => {}
_ = self.inner.runtime.sleep(Duration::from_millis(50)).fuse() => {}
}
}
}
impl<I> Drop for DataChannelImpl<I>
where
I: Interceptor,
{
fn drop(&mut self) {
if let Some(mut data_channels) = self.inner.data_channel_events_tx.try_lock() {
data_channels.remove(&self.id);
}
}
}
#[async_trait::async_trait]
impl<I> DataChannel for DataChannelImpl<I>
where
I: Interceptor + 'static,
{
async fn label(&self) -> Result<String> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.label()
.to_owned())
}
async fn ordered(&self) -> Result<bool> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.ordered())
}
async fn max_packet_life_time(&self) -> Result<Option<u16>> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.max_packet_life_time())
}
async fn max_retransmits(&self) -> Result<Option<u16>> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.max_retransmits())
}
async fn protocol(&self) -> Result<String> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.protocol()
.to_owned())
}
async fn negotiated(&self) -> Result<bool> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.negotiated())
}
fn id(&self) -> RTCDataChannelId {
self.id
}
async fn ready_state(&self) -> Result<RTCDataChannelState> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.ready_state())
}
async fn buffered_amount_high_threshold(&self) -> Result<u32> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.buffered_amount_high_threshold())
}
async fn set_buffered_amount_high_threshold(&self, threshold: u32) -> Result<()> {
{
let mut peer_connection = self.inner.core.lock().await;
peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.set_buffered_amount_high_threshold(threshold);
}
self.inner.wake_writes().await;
Ok(())
}
async fn buffered_amount_low_threshold(&self) -> Result<u32> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.buffered_amount_low_threshold())
}
async fn outstanding_bytes(&self) -> Result<usize> {
let mut peer_connection = self.inner.core.lock().await;
Ok(peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.outstanding_bytes())
}
async fn set_buffered_amount_low_threshold(&self, threshold: u32) -> Result<()> {
{
let mut peer_connection = self.inner.core.lock().await;
peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.set_buffered_amount_low_threshold(threshold);
}
self.inner.wake_writes().await;
Ok(())
}
async fn send(&self, data: BytesMut) -> Result<()> {
self.writable().await?;
{
let mut peer_connection = self.inner.core.lock().await;
let mut dc = peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?;
dc.send(data)?;
}
self.inner.wake_writes().await;
Ok(())
}
async fn send_text(&self, text: &str) -> Result<()> {
self.writable().await?;
{
let mut peer_connection = self.inner.core.lock().await;
let mut dc = peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?;
dc.send_text(text)?;
}
self.inner.wake_writes().await;
Ok(())
}
async fn writable(&self) -> Result<()> {
if self.inner.closing.load(Ordering::Acquire) {
return Err(Error::ErrDataChannelClosed);
}
let limit = self.inner.data_channel_send_buffer_limit;
if limit == usize::MAX {
return Ok(());
}
loop {
if self.inner.closing.load(Ordering::Acquire) {
return Err(Error::ErrDataChannelClosed);
}
{
let mut peer_connection = self.inner.core.lock().await;
let outstanding = peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.outstanding_bytes();
if outstanding < limit {
return Ok(());
}
} self.await_send_capacity().await;
}
}
async fn try_send(&self, data: BytesMut) -> Result<()> {
if self.inner.closing.load(Ordering::Acquire) {
return Err(Error::ErrDataChannelClosed);
}
let limit = self.inner.data_channel_send_buffer_limit;
{
let mut peer_connection = self.inner.core.lock().await;
let mut dc = peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?;
let outstanding = dc.outstanding_bytes();
if outstanding != 0 && outstanding.saturating_add(data.len()) > limit {
return Err(Error::ErrSendBufferFull);
}
dc.send(data)?;
}
self.inner.wake_writes().await;
Ok(())
}
async fn try_send_text(&self, text: &str) -> Result<()> {
if self.inner.closing.load(Ordering::Acquire) {
return Err(Error::ErrDataChannelClosed);
}
let limit = self.inner.data_channel_send_buffer_limit;
{
let mut peer_connection = self.inner.core.lock().await;
let mut dc = peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?;
let outstanding = dc.outstanding_bytes();
if outstanding != 0 && outstanding.saturating_add(text.len()) > limit {
return Err(Error::ErrSendBufferFull);
}
dc.send_text(text)?;
}
self.inner.wake_writes().await;
Ok(())
}
async fn poll(&self) -> Option<DataChannelEvent> {
self.evt_rx.lock().await.recv().await
}
async fn close(&self) -> Result<()> {
{
let mut peer_connection = self.inner.core.lock().await;
peer_connection
.data_channel(self.id)
.ok_or(Error::ErrDataChannelClosed)?
.close()?;
}
self.inner.wake_writes().await;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::executor::block_on;
use std::sync::atomic::AtomicUsize;
#[derive(Default)]
struct DefaultDataChannel {
sends: AtomicUsize,
text_sends: AtomicUsize,
}
#[async_trait::async_trait]
impl DataChannel for DefaultDataChannel {
async fn label(&self) -> Result<String> {
unimplemented!()
}
async fn ordered(&self) -> Result<bool> {
unimplemented!()
}
async fn max_packet_life_time(&self) -> Result<Option<u16>> {
unimplemented!()
}
async fn max_retransmits(&self) -> Result<Option<u16>> {
unimplemented!()
}
async fn protocol(&self) -> Result<String> {
unimplemented!()
}
async fn negotiated(&self) -> Result<bool> {
unimplemented!()
}
fn id(&self) -> RTCDataChannelId {
unimplemented!()
}
async fn ready_state(&self) -> Result<RTCDataChannelState> {
unimplemented!()
}
async fn buffered_amount_high_threshold(&self) -> Result<u32> {
unimplemented!()
}
async fn set_buffered_amount_high_threshold(&self, _threshold: u32) -> Result<()> {
unimplemented!()
}
async fn buffered_amount_low_threshold(&self) -> Result<u32> {
unimplemented!()
}
async fn set_buffered_amount_low_threshold(&self, _threshold: u32) -> Result<()> {
unimplemented!()
}
async fn send(&self, _data: BytesMut) -> Result<()> {
self.sends.fetch_add(1, Ordering::Relaxed);
Ok(())
}
async fn send_text(&self, _text: &str) -> Result<()> {
self.text_sends.fetch_add(1, Ordering::Relaxed);
Ok(())
}
async fn poll(&self) -> Option<DataChannelEvent> {
unimplemented!()
}
async fn close(&self) -> Result<()> {
unimplemented!()
}
}
#[test]
fn outstanding_bytes_trait_default_is_zero() {
let dc = DefaultDataChannel::default();
let n =
block_on(dc.outstanding_bytes()).expect("default outstanding_bytes() must return Ok");
assert_eq!(
n, 0,
"the DataChannel::outstanding_bytes default must report 0 outstanding bytes"
);
}
#[test]
fn writable_trait_default_is_ok() {
let dc = DefaultDataChannel::default();
block_on(dc.writable()).expect("the DataChannel::writable default must resolve Ok(())");
}
#[test]
fn try_send_trait_default_delegates_to_send() {
let dc = DefaultDataChannel::default();
block_on(dc.try_send(BytesMut::new())).expect("default try_send() must return Ok");
assert_eq!(
dc.sends.load(Ordering::Relaxed),
1,
"the DataChannel::try_send default must delegate to send()"
);
assert_eq!(dc.text_sends.load(Ordering::Relaxed), 0);
}
#[test]
fn try_send_text_trait_default_delegates_to_send_text() {
let dc = DefaultDataChannel::default();
block_on(dc.try_send_text("x")).expect("default try_send_text() must return Ok");
assert_eq!(
dc.text_sends.load(Ordering::Relaxed),
1,
"the DataChannel::try_send_text default must delegate to send_text()"
);
assert_eq!(dc.sends.load(Ordering::Relaxed), 0);
}
#[test]
fn drop_removes_event_sender() {
use crate::peer_connection::new_test_peer_connection;
use crate::runtime::{channel, default_runtime};
let rt = default_runtime().expect("test requires a runtime feature");
rt.block_on(Box::pin(async {
let (inner, _driver_event_rx) = new_test_peer_connection().await;
let channel_id = 0;
let (evt_tx, evt_rx) = channel::<DataChannelEvent>(1);
inner
.data_channel_events_tx
.lock()
.await
.insert(channel_id, evt_tx);
assert!(
inner
.data_channel_events_tx
.lock()
.await
.contains_key(&channel_id)
);
let dc = DataChannelImpl::new(channel_id, inner.clone(), evt_rx);
drop(dc);
assert!(
!inner
.data_channel_events_tx
.lock()
.await
.contains_key(&channel_id)
);
}));
}
}