use std::cmp::Ordering;
use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap};
use std::io;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
use std::sync::Arc;
use std::time::Duration;
use parking_lot::Mutex;
use rand::rngs::StdRng;
use rand::SeedableRng;
use tokio::sync::Notify;
use tokio::task::JoinHandle;
use tokio::time::Instant;
use super::{stream_seed, Netem};
use crate::Transport;
const TIMER_RESOLUTION: Duration = Duration::from_millis(1);
#[derive(Debug)]
struct Pending {
due: Instant,
seq: u64,
destination: SocketAddr,
bytes: Vec<u8>,
}
impl Ord for Pending {
fn cmp(&self, other: &Pending) -> Ordering {
self.due.cmp(&other.due).then(self.seq.cmp(&other.seq))
}
}
impl PartialOrd for Pending {
fn partial_cmp(&self, other: &Pending) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl PartialEq for Pending {
fn eq(&self, other: &Pending) -> bool {
self.cmp(other) == Ordering::Equal
}
}
impl Eq for Pending {}
#[derive(Default)]
struct InFlight {
queue: Mutex<BinaryHeap<Reverse<Pending>>>,
wake: Notify,
}
#[derive(Clone, Debug, Default)]
pub struct Impairments {
offered: Arc<AtomicU64>,
dropped: Arc<AtomicU64>,
}
impl Impairments {
pub fn offered(&self) -> u64 {
self.offered.load(AtomicOrdering::Relaxed)
}
pub fn dropped(&self) -> u64 {
self.dropped.load(AtomicOrdering::Relaxed)
}
pub fn loss_fraction(&self) -> f64 {
let offered = self.offered();
if offered == 0 {
0.0
} else {
self.dropped() as f64 / offered as f64
}
}
}
pub struct NetemTransport<T> {
inner: Arc<T>,
netem: Netem,
local: SocketAddr,
streams: Mutex<HashMap<SocketAddr, StdRng>>,
in_flight: Arc<InFlight>,
seq: AtomicU64,
impairments: Impairments,
pump: JoinHandle<()>,
}
impl<T: Transport> NetemTransport<T> {
pub fn new(inner: Arc<T>, netem: Netem) -> NetemTransport<T> {
let local = inner
.local_addr()
.expect("a netem-wrapped transport must already be bound");
let in_flight = Arc::new(InFlight::default());
let pump = tokio::spawn(pump(Arc::clone(&inner), Arc::clone(&in_flight)));
NetemTransport {
inner,
netem,
local,
streams: Mutex::new(HashMap::new()),
in_flight,
seq: AtomicU64::new(0),
impairments: Impairments::default(),
pump,
}
}
pub fn impairments(&self) -> Impairments {
self.impairments.clone()
}
}
impl<T> Drop for NetemTransport<T> {
fn drop(&mut self) {
self.pump.abort();
}
}
#[async_trait::async_trait]
impl<T: Transport> Transport for NetemTransport<T> {
async fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
self.inner.recv_from(buf).await
}
async fn send_to(&self, buf: &[u8], destination: &SocketAddr) -> io::Result<usize> {
let link = self.netem.link_to(destination);
let delay = {
let mut streams = self.streams.lock();
let stream = streams.entry(*destination).or_insert_with(|| {
StdRng::seed_from_u64(stream_seed(self.netem.seed, self.local, *destination))
});
link.draw(stream)
};
self.impairments
.offered
.fetch_add(1, AtomicOrdering::Relaxed);
let Some(delay) = delay else {
self.impairments
.dropped
.fetch_add(1, AtomicOrdering::Relaxed);
return Ok(buf.len());
};
self.in_flight.queue.lock().push(Reverse(Pending {
due: Instant::now() + delay,
seq: self.seq.fetch_add(1, AtomicOrdering::Relaxed),
destination: *destination,
bytes: buf.to_vec(),
}));
self.in_flight.wake.notify_one();
Ok(buf.len())
}
fn local_addr(&self) -> io::Result<SocketAddr> {
self.inner.local_addr()
}
}
enum Step {
Deliver(Pending),
WaitUntil(Instant),
Idle,
}
async fn pump<T: Transport>(inner: Arc<T>, in_flight: Arc<InFlight>) {
loop {
let step = {
let mut queue = in_flight.queue.lock();
match queue.peek().map(|Reverse(head)| head.due) {
Some(due) if due <= Instant::now() => {
Step::Deliver(queue.pop().expect("just peeked").0)
}
Some(due) => Step::WaitUntil(due),
None => Step::Idle,
}
};
match step {
Step::Deliver(datagram) => {
let _ = inner.send_to(&datagram.bytes, &datagram.destination).await;
}
Step::Idle => in_flight.wake.notified().await,
Step::WaitUntil(due) => wait_until(due, &in_flight.wake).await,
}
}
}
async fn wait_until(due: Instant, wake: &Notify) {
if let Some(coarse) = due.checked_sub(TIMER_RESOLUTION) {
if coarse_wait_worthwhile(coarse, Instant::now()) {
tokio::select! {
_ = tokio::time::sleep_until(coarse) => {}
_ = wake.notified() => return,
}
}
}
while still_waiting(Instant::now(), due) {
tokio::task::yield_now().await;
}
}
fn coarse_wait_worthwhile(coarse: Instant, now: Instant) -> bool {
coarse > now
}
fn still_waiting(now: Instant, due: Instant) -> bool {
now < due
}
#[cfg(test)]
mod tests;