use std::sync::{Arc, Mutex, RwLock};
use thiserror::Error;
use tokio::time::{Duration, Instant};
use crate::{frame::PingFrame, packet::PacketContent};
#[derive(Debug, Error)]
#[error("Path has been idle for too long")]
pub struct TimeOut;
#[derive(Debug)]
pub struct IdleConfig {
max_idle_timeout: Duration,
defer_idle_timeout: Duration,
heartbeat_interval: Duration,
}
impl IdleConfig {
fn suitable_heartbeat_interval(max_idle_timeout: Duration) -> Duration {
if max_idle_timeout == Duration::ZERO {
Duration::from_secs(30)
} else {
(max_idle_timeout / 2)
.max(Duration::from_secs(1))
.min(Duration::from_secs(30))
}
}
pub fn new(max_idle_timeout: Duration, defer_idle_timeout: Duration) -> Self {
let heartbeat_interval = Self::suitable_heartbeat_interval(max_idle_timeout);
Self {
max_idle_timeout,
defer_idle_timeout,
heartbeat_interval,
}
}
pub fn negotiate_max_idle_timeout(&mut self, max_idle_timeout: Duration) {
match (self.max_idle_timeout, max_idle_timeout) {
(_, Duration::ZERO) => (),
(Duration::ZERO, remote) => self.max_idle_timeout = remote,
(local, remote) => self.max_idle_timeout = local.min(remote),
}
self.heartbeat_interval = Self::suitable_heartbeat_interval(self.max_idle_timeout);
}
pub fn set_heartbeat_interval(&mut self, interval: Duration) {
self.heartbeat_interval = interval;
}
}
#[derive(Debug, Clone)]
pub struct ArcIdleConfig(Arc<RwLock<IdleConfig>>);
impl ArcIdleConfig {
pub fn new(max_idle_timeout: Duration, defer_idle_timeout: Duration) -> Self {
ArcIdleConfig(Arc::new(RwLock::new(IdleConfig::new(
max_idle_timeout,
defer_idle_timeout,
))))
}
pub fn negotiate_max_idle_timeout(&self, max_idle_timeout: Duration) {
self.0
.write()
.unwrap()
.negotiate_max_idle_timeout(max_idle_timeout);
}
pub fn set_heartbeat_interval(&self, interval: Duration) {
self.0.write().unwrap().set_heartbeat_interval(interval);
}
pub fn timer(&self) -> ArcIdleTimer {
ArcIdleTimer(Arc::new(Mutex::new(IdleTimer {
idle_config: self.clone(),
heartbeat_times: 0,
last_effective_comm: None,
idle_begin_at: Some(Instant::now()),
})))
}
fn defer_idle_timeout(&self) -> Duration {
self.0.read().unwrap().defer_idle_timeout
}
fn heartbeat_interval(&self) -> Duration {
self.0.read().unwrap().heartbeat_interval
}
fn timeout_after(&self, idle_at: Instant) -> bool {
let max_idle_timeout = self.0.read().unwrap().max_idle_timeout;
max_idle_timeout != Duration::ZERO && idle_at.elapsed() > max_idle_timeout
}
}
#[derive(Debug)]
pub struct IdleTimer {
idle_config: ArcIdleConfig,
heartbeat_times: u32,
last_effective_comm: Option<Instant>,
idle_begin_at: Option<Instant>,
}
impl IdleTimer {
pub fn on_sent(&mut self, _packet_content: PacketContent) {
}
pub fn on_rcvd(&mut self, packet_content: PacketContent) {
if packet_content == PacketContent::EffectivePayload {
self.last_effective_comm = Some(Instant::now());
self.heartbeat_times = 0;
self.idle_begin_at = None;
}
if self.idle_begin_at.is_some() {
self.idle_begin_at = Some(Instant::now());
}
}
pub fn health(&mut self) -> Result<Option<PingFrame>, TimeOut> {
if let Some(t) = self.last_effective_comm {
let elapsed = t.elapsed();
if elapsed > self.idle_config.defer_idle_timeout() {
if self.idle_begin_at.is_none() {
self.idle_begin_at = Some(Instant::now());
return Ok(Some(PingFrame)); }
} else if elapsed > self.idle_config.heartbeat_interval() * (self.heartbeat_times + 1) {
self.heartbeat_times += 1;
return Ok(Some(PingFrame));
}
}
if self
.idle_begin_at
.is_some_and(|t| self.idle_config.timeout_after(t))
{
return Err(TimeOut);
}
Ok(None)
}
}
#[derive(Debug, Clone)]
pub struct ArcIdleTimer(Arc<Mutex<IdleTimer>>);
impl ArcIdleTimer {
pub fn on_sent(&self, packet_content: PacketContent) {
self.0.lock().unwrap().on_sent(packet_content);
}
pub fn on_rcvd(&self, packet_content: PacketContent) {
self.0.lock().unwrap().on_rcvd(packet_content);
}
pub fn health(&self) -> Result<Option<PingFrame>, TimeOut> {
self.0.lock().unwrap().health()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sent_effective_payload_does_not_extend_idle_timeout() {
let mut timer = IdleTimer {
idle_config: ArcIdleConfig::new(Duration::from_millis(10), Duration::ZERO),
heartbeat_times: 0,
last_effective_comm: None,
idle_begin_at: None,
};
timer.on_rcvd(PacketContent::EffectivePayload);
assert!(timer.health().expect("timer should be healthy").is_some());
std::thread::sleep(Duration::from_millis(6));
timer.on_sent(PacketContent::EffectivePayload);
std::thread::sleep(Duration::from_millis(6));
assert!(timer.health().is_err());
}
#[test]
fn new_timer_times_out_without_receive() {
let timer = ArcIdleConfig::new(Duration::from_millis(10), Duration::ZERO).timer();
std::thread::sleep(Duration::from_millis(20));
assert!(timer.health().is_err());
}
#[test]
fn zero_max_idle_timeout_disables_new_timer_timeout() {
let timer = ArcIdleConfig::new(Duration::ZERO, Duration::ZERO).timer();
std::thread::sleep(Duration::from_millis(20));
assert!(timer.health().is_ok());
}
#[test]
fn effective_receive_clears_initial_idle_deadline() {
let timer = ArcIdleConfig::new(Duration::from_millis(10), Duration::ZERO).timer();
timer.on_rcvd(PacketContent::EffectivePayload);
std::thread::sleep(Duration::from_millis(20));
assert!(timer.health().is_ok());
}
}