use std::time::{Duration, Instant};
pub mod constants {
use std::time::Duration;
pub const INITIAL_RTO: Duration = Duration::from_millis(1000);
pub const MIN_RTO: Duration = Duration::from_millis(100);
pub const MAX_RTO: Duration = Duration::from_millis(60000);
pub const SRTT_ALPHA: f64 = 0.125;
pub const RTTVAR_BETA: f64 = 0.25;
pub const RTO_K: f64 = 4.0;
pub const MIN_RTO_GRANULARITY_MS: f64 = 100.0;
}
#[derive(Debug, Clone)]
pub struct RttEstimator {
srtt: f64,
rttvar: f64,
rto: Duration,
initialized: bool,
}
impl Default for RttEstimator {
fn default() -> Self {
Self::new()
}
}
impl RttEstimator {
pub fn new() -> Self {
Self {
srtt: 0.0,
rttvar: 0.0,
rto: constants::INITIAL_RTO,
initialized: false,
}
}
pub fn update(&mut self, sample: Duration) {
let sample_ms = sample.as_secs_f64() * 1000.0;
if !self.initialized {
self.srtt = sample_ms;
self.rttvar = sample_ms / 2.0;
self.initialized = true;
} else {
self.rttvar = (1.0 - constants::RTTVAR_BETA) * self.rttvar
+ constants::RTTVAR_BETA * (self.srtt - sample_ms).abs();
self.srtt =
(1.0 - constants::SRTT_ALPHA) * self.srtt + constants::SRTT_ALPHA * sample_ms;
}
let rto_ms =
self.srtt + f64::max(constants::MIN_RTO_GRANULARITY_MS, constants::RTO_K * self.rttvar);
let rto_ms = rto_ms.clamp(
constants::MIN_RTO.as_millis() as f64,
constants::MAX_RTO.as_millis() as f64,
);
self.rto = Duration::from_millis(rto_ms as u64);
}
pub fn srtt(&self) -> Duration {
Duration::from_secs_f64(self.srtt / 1000.0)
}
pub fn srtt_ms(&self) -> f64 {
self.srtt
}
pub fn rttvar(&self) -> Duration {
Duration::from_secs_f64(self.rttvar / 1000.0)
}
pub fn rto(&self) -> Duration {
self.rto
}
pub fn is_initialized(&self) -> bool {
self.initialized
}
pub fn backoff(&mut self) -> Duration {
let new_rto_ms = (self.rto.as_millis() as u64).saturating_mul(2);
self.rto = Duration::from_millis(new_rto_ms).min(constants::MAX_RTO);
self.rto
}
pub fn reset_backoff(&mut self) {
if self.initialized {
let rto_ms = self.srtt
+ f64::max(constants::MIN_RTO_GRANULARITY_MS, constants::RTO_K * self.rttvar);
let rto_ms = rto_ms.clamp(
constants::MIN_RTO.as_millis() as f64,
constants::MAX_RTO.as_millis() as f64,
);
self.rto = Duration::from_millis(rto_ms as u64);
} else {
self.rto = constants::INITIAL_RTO;
}
}
}
#[derive(Debug, Clone)]
pub struct TimestampTracker {
session_start: Instant,
last_peer_timestamp: u32,
pending_timestamp: Option<u32>,
pending_send_time: Option<Instant>,
}
impl TimestampTracker {
pub fn new() -> Self {
Self {
session_start: Instant::now(),
last_peer_timestamp: 0,
pending_timestamp: None,
pending_send_time: None,
}
}
pub fn with_start(start: Instant) -> Self {
Self {
session_start: start,
last_peer_timestamp: 0,
pending_timestamp: None,
pending_send_time: None,
}
}
pub fn now(&self) -> u32 {
self.session_start.elapsed().as_millis() as u32
}
pub fn timestamp_echo(&self) -> u32 {
self.last_peer_timestamp
}
pub fn on_send(&mut self, timestamp: u32) {
self.pending_timestamp = Some(timestamp);
self.pending_send_time = Some(Instant::now());
}
pub fn on_receive(&mut self, peer_timestamp: u32, echo: u32) -> Option<Duration> {
self.last_peer_timestamp = peer_timestamp;
if let (Some(pending), Some(send_time)) = (self.pending_timestamp, self.pending_send_time)
&& echo == pending
{
let rtt = send_time.elapsed();
self.pending_timestamp = None;
self.pending_send_time = None;
return Some(rtt);
}
None
}
pub fn clear_pending(&mut self) {
self.pending_timestamp = None;
self.pending_send_time = None;
}
}
impl Default for TimestampTracker {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rtt_estimator_initial() {
let estimator = RttEstimator::new();
assert!(!estimator.is_initialized());
assert_eq!(estimator.rto(), constants::INITIAL_RTO);
}
#[test]
fn test_rtt_estimator_first_sample() {
let mut estimator = RttEstimator::new();
estimator.update(Duration::from_millis(100));
assert!(estimator.is_initialized());
assert!((estimator.srtt_ms() - 100.0).abs() < 0.01);
assert!((estimator.rttvar - 50.0).abs() < 0.01); }
#[test]
fn test_rtt_estimator_multiple_samples() {
let mut estimator = RttEstimator::new();
estimator.update(Duration::from_millis(100));
let srtt1 = estimator.srtt_ms();
estimator.update(Duration::from_millis(120));
let srtt2 = estimator.srtt_ms();
assert!(srtt2 > srtt1);
assert!(srtt2 < 120.0);
}
#[test]
fn test_rtt_estimator_backoff() {
let mut estimator = RttEstimator::new();
estimator.update(Duration::from_millis(100));
let rto1 = estimator.rto();
let rto2 = estimator.backoff();
assert!(rto2 > rto1);
assert!(rto2 <= constants::MAX_RTO);
}
#[test]
fn test_rtt_estimator_max_rto() {
let mut estimator = RttEstimator::new();
estimator.update(Duration::from_millis(100));
for _ in 0..20 {
estimator.backoff();
}
assert_eq!(estimator.rto(), constants::MAX_RTO);
}
#[test]
fn test_rtt_estimator_min_rto() {
let mut estimator = RttEstimator::new();
estimator.update(Duration::from_micros(100));
assert!(estimator.rto() >= constants::MIN_RTO);
}
#[test]
fn test_timestamp_tracker_echo() {
let start = Instant::now();
let mut tracker = TimestampTracker::with_start(start);
tracker.on_send(1000);
std::thread::sleep(Duration::from_millis(10));
let rtt = tracker.on_receive(2000, 1000);
assert!(rtt.is_some());
let rtt = rtt.unwrap();
assert!(rtt >= Duration::from_millis(10));
}
#[test]
fn test_timestamp_tracker_no_match() {
let start = Instant::now();
let mut tracker = TimestampTracker::with_start(start);
tracker.on_send(1000);
let rtt = tracker.on_receive(2000, 999);
assert!(rtt.is_none());
assert!(tracker.pending_timestamp.is_some());
}
#[test]
fn test_timestamp_tracker_peer_timestamp() {
let mut tracker = TimestampTracker::new();
assert_eq!(tracker.timestamp_echo(), 0);
tracker.on_receive(5000, 0);
assert_eq!(tracker.timestamp_echo(), 5000);
}
}