use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use livekit::prelude::{DataTrackFrame, LocalDataTrack, LocalParticipant, PublishError};
use tokio::runtime::Handle;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use tracing::{debug, error, warn};
use crate::{ChannelId, Metadata};
const FRAME_HEADER_SIZE: usize = 8;
pub(super) const OVERSIZED_WARN_INTERVAL: Duration = Duration::from_secs(30);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct OversizedDropReport {
pub dropped_since_last: u64,
pub size_limit: usize,
}
pub(crate) struct DataTrack {
track: Arc<OnceLock<LocalDataTrack>>,
close: CancellationToken,
task: Option<JoinHandle<()>>,
sequence: AtomicU32,
drop_throttler: parking_lot::Mutex<crate::throttler::Throttler>,
max_message_size: usize,
oversized_dropped: AtomicU64,
oversized_throttler: parking_lot::Mutex<crate::throttler::Throttler>,
}
impl DataTrack {
pub fn publish(
runtime: &Handle,
local_participant: LocalParticipant,
channel_id: ChannelId,
session_cancel: CancellationToken,
max_message_size: usize,
) -> Self {
let track = Arc::new(OnceLock::new());
let track_clone = Arc::clone(&track);
let close = CancellationToken::new();
let close_clone = close.clone();
let name = format!("data-ch-{}", u64::from(channel_id));
let task = runtime.spawn(async move {
const INITIAL_BACKOFF: Duration = Duration::from_millis(100);
const MAX_BACKOFF: Duration = Duration::from_secs(3);
let mut backoff = INITIAL_BACKOFF;
loop {
if close_clone.is_cancelled() {
return;
}
let result = tokio::select! {
() = session_cancel.cancelled() => return,
result = local_participant.publish_data_track(name.clone()) => result,
};
match result {
Ok(published) => {
track_clone.set(published).ok();
debug!("data track {name} published");
return;
}
Err(PublishError::DuplicateName) => {
debug!(
"data track {name} still being unpublished at SFU, \
retrying in {backoff:?}"
);
}
Err(e) => {
error!(
"failed to publish data track {name}: {e:?}, \
retrying in {backoff:?}"
);
}
}
tokio::select! {
() = close_clone.cancelled() => return,
() = session_cancel.cancelled() => return,
() = tokio::time::sleep(backoff) => {}
}
backoff = (backoff * 2).min(MAX_BACKOFF);
}
});
Self {
track,
close,
task: Some(task),
sequence: AtomicU32::new(0),
drop_throttler: parking_lot::Mutex::new(crate::throttler::Throttler::new(
Duration::from_secs(30),
)),
max_message_size,
oversized_dropped: AtomicU64::new(0),
oversized_throttler: parking_lot::Mutex::new(crate::throttler::Throttler::new(
OVERSIZED_WARN_INTERVAL,
)),
}
}
pub fn log(
&self,
channel_id: ChannelId,
msg: &[u8],
metadata: &Metadata,
) -> Result<(), OversizedDropReport> {
if msg.len() > self.max_message_size {
if self.oversized_throttler.lock().try_acquire() {
let dropped = 1 + self.oversized_dropped.swap(0, Ordering::Relaxed);
warn!(
"dropping {}-byte message on channel {channel_id:?}: exceeds \
data-track limit of {} bytes ({dropped} dropped since last warning)",
msg.len(),
self.max_message_size
);
return Err(OversizedDropReport {
dropped_since_last: dropped,
size_limit: self.max_message_size,
});
}
self.oversized_dropped.fetch_add(1, Ordering::Relaxed);
return Ok(());
}
let Some(track) = self.track.get() else {
if self.drop_throttler.lock().try_acquire() {
debug!("data track not ready, dropping message for channel {channel_id:?}");
}
return Ok(());
};
let seq = self.sequence.fetch_add(1, Ordering::Relaxed);
let mut payload = Vec::with_capacity(FRAME_HEADER_SIZE + msg.len());
payload.extend_from_slice(&0u16.to_le_bytes());
payload.extend_from_slice(&(FRAME_HEADER_SIZE as u16).to_le_bytes()); payload.extend_from_slice(&seq.to_le_bytes());
payload.extend_from_slice(msg);
let frame = DataTrackFrame::new(payload).with_user_timestamp(metadata.log_time);
if let Err(e) = track.try_push(frame)
&& self.drop_throttler.lock().try_acquire()
{
debug!("data track message dropped for channel {channel_id:?}: {e:?}");
}
Ok(())
}
#[cfg(test)]
pub fn oversized_dropped(&self) -> u64 {
self.oversized_dropped.load(Ordering::Relaxed)
}
pub async fn close(&mut self) {
self.close.cancel();
if let Some(task) = self.task.take() {
_ = task.await;
}
if let Some(track) = self.track.get() {
debug!("unpublishing data track {}", track.info().name());
track.unpublish();
}
}
}
impl Drop for DataTrack {
fn drop(&mut self) {
self.close.cancel();
}
}
#[cfg(test)]
mod tests {
use super::*;
impl DataTrack {
fn for_test(max_message_size: usize) -> Self {
Self {
track: Arc::new(OnceLock::new()),
close: CancellationToken::new(),
task: None,
sequence: AtomicU32::new(0),
drop_throttler: parking_lot::Mutex::new(crate::throttler::Throttler::new(
Duration::from_secs(30),
)),
max_message_size,
oversized_dropped: AtomicU64::new(0),
oversized_throttler: parking_lot::Mutex::new(crate::throttler::Throttler::new(
OVERSIZED_WARN_INTERVAL,
)),
}
}
}
#[test]
fn drops_oversized_message_and_counts_it() {
let track = DataTrack::for_test(16);
let channel_id = ChannelId::new(1);
let metadata = Metadata::default();
let report = track.log(channel_id, &[0u8; 17], &metadata);
assert_eq!(
report,
Err(OversizedDropReport {
dropped_since_last: 1,
size_limit: 16,
})
);
assert_eq!(track.log(channel_id, &[0u8; 17], &metadata), Ok(()));
assert_eq!(track.oversized_dropped(), 1);
assert_eq!(track.log(channel_id, &[0u8; 16], &metadata), Ok(()));
assert_eq!(track.oversized_dropped(), 1);
}
}