use std::os::fd::{AsFd, AsRawFd, BorrowedFd};
use std::time::Duration;
use tokio::io::unix::AsyncFd;
use crate::afpacket::tx::Injector;
use crate::error::Error;
pub struct AsyncInjector {
inner: AsyncFd<Injector>,
}
impl AsyncInjector {
pub fn new(tx: Injector) -> Result<Self, Error> {
let fd = AsyncFd::with_interest(tx, tokio::io::Interest::WRITABLE).map_err(Error::Io)?;
Ok(Self { inner: fd })
}
pub fn open(interface: &str) -> Result<Self, Error> {
Self::new(Injector::open(interface)?)
}
pub async fn send(&mut self, data: &[u8]) -> Result<(), Error> {
let cap = self.inner.get_ref().frame_capacity();
if data.len() > cap {
return Err(Error::Config(format!(
"packet length {} exceeds TX frame capacity {}",
data.len(),
cap
)));
}
loop {
if let Some(mut slot) = self.inner.get_mut().allocate(data.len()) {
slot.data_mut()[..data.len()].copy_from_slice(data);
slot.set_len(data.len());
slot.send();
return Ok(());
}
let mut guard = self.inner.writable_mut().await.map_err(Error::Io)?;
guard.clear_ready();
drop(guard);
}
}
pub async fn flush(&mut self) -> Result<usize, Error> {
self.inner.get_mut().flush()
}
pub async fn wait_drained(&mut self, timeout: Duration) -> Result<(), Error> {
let deadline = tokio::time::Instant::now() + timeout;
loop {
if self.inner.get_ref().pending_count() == 0 {
return Ok(());
}
let remaining = match deadline.checked_duration_since(tokio::time::Instant::now()) {
Some(r) => r,
None => {
return Err(Error::Io(std::io::Error::from(
std::io::ErrorKind::TimedOut,
)));
}
};
let slice = remaining.min(Duration::from_millis(10));
tokio::select! {
ready = self.inner.writable_mut() => {
let mut guard = ready.map_err(Error::Io)?;
guard.clear_ready();
}
_ = tokio::time::sleep(slice) => {}
}
}
}
pub async fn send_stream<S, B>(
&mut self,
stream: S,
mut pacer: Option<TxPacer>,
) -> Result<usize, Error>
where
S: futures_core::Stream<Item = B>,
B: AsRef<[u8]>,
{
let mut stream = std::pin::pin!(stream);
let mut sent = 0usize;
while let Some(frame) = std::future::poll_fn(|cx| stream.as_mut().poll_next(cx)).await {
let bytes = frame.as_ref();
if let Some(p) = pacer.as_mut() {
let wait = p.acquire(bytes.len(), std::time::Instant::now());
if !wait.is_zero() {
tokio::time::sleep(wait).await;
}
}
self.send(bytes).await?;
self.flush().await?;
sent += 1;
}
Ok(sent)
}
pub fn read_tx_timestamp(&self) -> Option<crate::Timestamp> {
self.inner.get_ref().read_tx_timestamp()
}
pub fn get_ref(&self) -> &Injector {
self.inner.get_ref()
}
pub fn get_mut(&mut self) -> &mut Injector {
self.inner.get_mut()
}
pub fn into_inner(self) -> Injector {
self.inner.into_inner()
}
#[inline]
pub fn frame_capacity(&self) -> usize {
self.inner.get_ref().frame_capacity()
}
#[inline]
pub fn frame_count(&self) -> usize {
self.inner.get_ref().frame_count()
}
pub fn available_slots(&self) -> usize {
self.inner.get_ref().available_slots()
}
pub fn rejected_slots(&self) -> usize {
self.inner.get_ref().rejected_slots()
}
pub fn pending_count(&self) -> usize {
self.inner.get_ref().pending_count()
}
}
impl AsFd for AsyncInjector {
fn as_fd(&self) -> BorrowedFd<'_> {
self.inner.get_ref().as_fd()
}
}
impl AsRawFd for AsyncInjector {
fn as_raw_fd(&self) -> std::os::fd::RawFd {
self.inner.get_ref().as_raw_fd()
}
}
#[derive(Debug, Clone)]
pub struct TxPacer {
rate: f64,
burst: f64,
tokens: f64,
last: Option<std::time::Instant>,
cost_bits: bool,
}
impl TxPacer {
fn new(rate: f64, burst: f64, cost_bits: bool) -> Self {
let burst = burst.max(1.0);
Self {
rate: rate.max(f64::MIN_POSITIVE),
burst,
tokens: burst,
last: None,
cost_bits,
}
}
pub fn packets_per_second(pps: f64) -> Self {
Self::new(pps, pps, false)
}
pub fn bits_per_second(bps: f64) -> Self {
Self::new(bps, bps, true)
}
pub fn with_burst(mut self, burst: f64) -> Self {
self.burst = burst.max(1.0);
self.tokens = self.burst;
self
}
fn cost(&self, len: usize) -> f64 {
if self.cost_bits {
(len as f64) * 8.0
} else {
1.0
}
}
pub fn acquire(&mut self, len: usize, now: std::time::Instant) -> Duration {
if let Some(last) = self.last {
let elapsed = now.saturating_duration_since(last).as_secs_f64();
self.tokens = (self.tokens + elapsed * self.rate).min(self.burst);
}
self.last = Some(now);
let cost = self.cost(len);
if self.tokens >= cost {
self.tokens -= cost;
Duration::ZERO
} else {
let deficit = cost - self.tokens;
self.tokens = 0.0;
Duration::from_secs_f64(deficit / self.rate)
}
}
}
#[cfg(test)]
mod pacer_tests {
use super::TxPacer;
use std::time::{Duration, Instant};
#[test]
fn first_burst_is_free_then_paces() {
let t0 = Instant::now();
let mut p = TxPacer::packets_per_second(100.0).with_burst(1.0);
assert_eq!(
p.acquire(64, t0),
Duration::ZERO,
"first frame within burst"
);
let wait = p.acquire(64, t0);
assert!(wait > Duration::ZERO, "second back-to-back frame must wait");
assert!(
(wait.as_secs_f64() - 0.01).abs() < 1e-3,
"expected ~10ms, got {wait:?}"
);
}
#[test]
fn refills_over_time_no_wait_after_interval() {
let t0 = Instant::now();
let mut p = TxPacer::packets_per_second(100.0).with_burst(1.0);
let _ = p.acquire(64, t0);
let t1 = t0 + Duration::from_millis(20);
assert_eq!(p.acquire(64, t1), Duration::ZERO);
}
#[test]
fn bits_per_second_costs_scale_with_frame_size() {
let t0 = Instant::now();
let mut p = TxPacer::bits_per_second(8000.0).with_burst(8000.0);
assert_eq!(p.acquire(1000, t0), Duration::ZERO);
let wait = p.acquire(1000, t0);
assert!(
(wait.as_secs_f64() - 1.0).abs() < 1e-2,
"expected ~1s, got {wait:?}"
);
}
}