use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct PendingAck {
pub version: u64,
pub sent_at: Instant,
pub retransmit_count: u32,
pub rto: Duration,
}
impl PendingAck {
pub fn new(version: u64, rto: Duration) -> Self {
Self {
version,
sent_at: Instant::now(),
retransmit_count: 0,
rto,
}
}
pub fn needs_retransmit(&self) -> bool {
self.sent_at.elapsed() >= self.rto
}
pub fn retransmit(&mut self, backoff_multiplier: u32, max_rto: Duration) {
self.sent_at = Instant::now();
self.retransmit_count += 1;
self.rto = (self.rto * backoff_multiplier).min(max_rto);
}
pub fn time_until_retransmit(&self) -> Duration {
let elapsed = self.sent_at.elapsed();
if elapsed >= self.rto {
Duration::ZERO
} else {
self.rto - elapsed
}
}
}
pub const DEFAULT_INITIAL_RTO: Duration = Duration::from_millis(1000);
pub const DEFAULT_MIN_RTO: Duration = Duration::from_millis(100);
pub const DEFAULT_MAX_RTO: Duration = Duration::from_secs(60);
pub const DEFAULT_BACKOFF_MULTIPLIER: u32 = 2;
pub const DEFAULT_MAX_RETRANSMITS: u32 = 10;
#[derive(Debug)]
pub struct AckTracker {
pending: Vec<PendingAck>,
highest_acked: u64,
initial_rto: Duration,
min_rto: Duration,
max_rto: Duration,
backoff_multiplier: u32,
max_retransmits: u32,
srtt: Option<Duration>,
rttvar: Option<Duration>,
}
impl AckTracker {
pub fn new() -> Self {
Self {
pending: Vec::new(),
highest_acked: 0,
initial_rto: DEFAULT_INITIAL_RTO,
min_rto: DEFAULT_MIN_RTO,
max_rto: DEFAULT_MAX_RTO,
backoff_multiplier: DEFAULT_BACKOFF_MULTIPLIER,
max_retransmits: DEFAULT_MAX_RETRANSMITS,
srtt: None,
rttvar: None,
}
}
pub fn with_rto(
initial_rto: Duration,
min_rto: Duration,
max_rto: Duration,
backoff_multiplier: u32,
max_retransmits: u32,
) -> Self {
Self {
pending: Vec::new(),
highest_acked: 0,
initial_rto,
min_rto,
max_rto,
backoff_multiplier,
max_retransmits,
srtt: None,
rttvar: None,
}
}
pub fn register_sent(&mut self, version: u64) {
if self.pending.iter().any(|p| p.version == version) {
return;
}
let rto = self.current_rto();
self.pending.push(PendingAck::new(version, rto));
}
pub fn process_ack(&mut self, acked_version: u64) -> Option<Duration> {
if acked_version <= self.highest_acked {
return None;
}
self.highest_acked = acked_version;
let mut rtt_sample = None;
self.pending.retain(|pending| {
if pending.version <= acked_version {
if pending.retransmit_count == 0 && rtt_sample.is_none() {
rtt_sample = Some(pending.sent_at.elapsed());
}
false } else {
true }
});
if let Some(rtt) = rtt_sample {
self.update_rtt(rtt);
}
rtt_sample
}
fn update_rtt(&mut self, rtt: Duration) {
let rtt_secs = rtt.as_secs_f64();
match (self.srtt, self.rttvar) {
(None, None) => {
self.srtt = Some(rtt);
self.rttvar = Some(rtt / 2);
}
(Some(srtt), Some(rttvar)) => {
let srtt_secs = srtt.as_secs_f64();
let rttvar_secs = rttvar.as_secs_f64();
let new_rttvar =
0.75 * rttvar_secs + 0.25 * (srtt_secs - rtt_secs).abs();
let new_srtt = 0.875 * srtt_secs + 0.125 * rtt_secs;
self.srtt = Some(Duration::from_secs_f64(new_srtt));
self.rttvar = Some(Duration::from_secs_f64(new_rttvar));
}
_ => {}
}
}
pub fn current_rto(&self) -> Duration {
match (self.srtt, self.rttvar) {
(Some(srtt), Some(rttvar)) => {
let k = 4;
let g = Duration::from_millis(1);
let rto = srtt + (g.max(rttvar * k));
rto.clamp(self.min_rto, self.max_rto)
}
_ => self.initial_rto,
}
}
pub fn srtt(&self) -> Option<Duration> {
self.srtt
}
pub fn rttvar(&self) -> Option<Duration> {
self.rttvar
}
pub fn needs_retransmit(&self) -> impl Iterator<Item = u64> + '_ {
self.pending
.iter()
.filter(|p| p.needs_retransmit() && p.retransmit_count < self.max_retransmits)
.map(|p| p.version)
}
pub fn failed_versions(&self) -> impl Iterator<Item = u64> + '_ {
self.pending
.iter()
.filter(|p| p.retransmit_count >= self.max_retransmits)
.map(|p| p.version)
}
pub fn mark_retransmitted(&mut self, version: u64) {
if let Some(pending) = self.pending.iter_mut().find(|p| p.version == version) {
pending.retransmit(self.backoff_multiplier, self.max_rto);
}
}
pub fn has_pending(&self) -> bool {
!self.pending.is_empty()
}
pub fn pending_count(&self) -> usize {
self.pending.len()
}
pub fn highest_acked(&self) -> u64 {
self.highest_acked
}
pub fn time_until_retransmit(&self) -> Option<Duration> {
self.pending
.iter()
.filter(|p| p.retransmit_count < self.max_retransmits)
.map(|p| p.time_until_retransmit())
.min()
}
pub fn cancel(&mut self, version: u64) {
self.pending.retain(|p| p.version != version);
}
pub fn cancel_all(&mut self) {
self.pending.clear();
}
pub fn reset(&mut self) {
self.pending.clear();
self.highest_acked = 0;
self.srtt = None;
self.rttvar = None;
}
}
impl Default for AckTracker {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
#[test]
fn test_new_tracker() {
let tracker = AckTracker::new();
assert!(!tracker.has_pending());
assert_eq!(tracker.highest_acked(), 0);
assert_eq!(tracker.current_rto(), DEFAULT_INITIAL_RTO);
}
#[test]
fn test_register_sent() {
let mut tracker = AckTracker::new();
tracker.register_sent(1);
assert!(tracker.has_pending());
assert_eq!(tracker.pending_count(), 1);
tracker.register_sent(1);
assert_eq!(tracker.pending_count(), 1);
tracker.register_sent(2);
assert_eq!(tracker.pending_count(), 2);
}
#[test]
fn test_process_ack() {
let mut tracker = AckTracker::new();
tracker.register_sent(1);
tracker.register_sent(2);
tracker.register_sent(3);
tracker.process_ack(2);
assert_eq!(tracker.highest_acked(), 2);
assert_eq!(tracker.pending_count(), 1);
tracker.process_ack(1);
assert_eq!(tracker.highest_acked(), 2);
}
#[test]
fn test_rtt_sample() {
let mut tracker = AckTracker::new();
tracker.register_sent(1);
thread::sleep(Duration::from_millis(10));
let rtt = tracker.process_ack(1);
assert!(rtt.is_some());
assert!(rtt.unwrap() >= Duration::from_millis(10));
assert!(tracker.srtt().is_some());
assert!(tracker.rttvar().is_some());
}
#[test]
fn test_retransmit() {
let mut tracker = AckTracker::with_rto(
Duration::from_millis(10),
Duration::from_millis(10),
Duration::from_secs(1),
2,
3,
);
tracker.register_sent(1);
assert_eq!(tracker.needs_retransmit().count(), 0);
thread::sleep(Duration::from_millis(15));
let versions: Vec<_> = tracker.needs_retransmit().collect();
assert_eq!(versions, vec![1]);
tracker.mark_retransmitted(1);
assert_eq!(tracker.needs_retransmit().count(), 0);
}
#[test]
fn test_max_retransmits() {
let mut tracker = AckTracker::with_rto(
Duration::from_millis(1),
Duration::from_millis(1),
Duration::from_millis(10),
1, 2, );
tracker.register_sent(1);
thread::sleep(Duration::from_millis(5));
tracker.mark_retransmitted(1);
thread::sleep(Duration::from_millis(5));
tracker.mark_retransmitted(1);
thread::sleep(Duration::from_millis(5));
let failed: Vec<_> = tracker.failed_versions().collect();
assert_eq!(failed, vec![1]);
assert_eq!(tracker.needs_retransmit().count(), 0);
}
#[test]
fn test_cancel() {
let mut tracker = AckTracker::new();
tracker.register_sent(1);
tracker.register_sent(2);
tracker.register_sent(3);
tracker.cancel(2);
assert_eq!(tracker.pending_count(), 2);
tracker.cancel_all();
assert!(!tracker.has_pending());
}
#[test]
fn test_reset() {
let mut tracker = AckTracker::new();
tracker.register_sent(1);
tracker.process_ack(1);
tracker.reset();
assert!(!tracker.has_pending());
assert_eq!(tracker.highest_acked(), 0);
assert!(tracker.srtt().is_none());
}
#[test]
fn test_time_until_retransmit() {
let mut tracker = AckTracker::with_rto(
Duration::from_millis(100),
Duration::from_millis(100),
Duration::from_secs(1),
2,
10,
);
assert!(tracker.time_until_retransmit().is_none());
tracker.register_sent(1);
let time = tracker.time_until_retransmit();
assert!(time.is_some());
assert!(time.unwrap() <= Duration::from_millis(100));
}
}