use std::time::Duration;
use tokio::time::{Instant, Interval, MissedTickBehavior, interval_at};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HeartbeatAction {
Idle,
SendEmpty,
PeerTimedOut,
}
#[derive(Debug)]
pub struct Heartbeat {
timer: Option<Interval>,
send_after: Option<Duration>,
recv_timeout: Option<Duration>,
last_recv: Instant,
last_send: Instant,
}
impl Heartbeat {
pub fn new(send_after: Option<Duration>, recv_timeout: Option<Duration>) -> Self {
let check = match (send_after, recv_timeout) {
(Some(s), Some(r)) => Some(s.min(r / 2)),
(Some(s), None) => Some(s),
(None, Some(r)) => Some(r / 2),
(None, None) => None,
}
.map(|d| d.max(Duration::from_millis(1)));
let now = Instant::now();
let timer = check.map(|d| {
let mut iv = interval_at(now + d, d);
iv.set_missed_tick_behavior(MissedTickBehavior::Delay);
iv
});
Heartbeat {
timer,
send_after,
recv_timeout,
last_recv: now,
last_send: now,
}
}
pub fn record_recv(&mut self) {
self.last_recv = Instant::now();
}
pub fn record_send(&mut self) {
self.last_send = Instant::now();
}
pub async fn tick(&mut self) -> HeartbeatAction {
match &mut self.timer {
Some(iv) => {
iv.tick().await;
if let Some(to) = self.recv_timeout {
if self.last_recv.elapsed() >= to {
return HeartbeatAction::PeerTimedOut;
}
}
if let Some(sa) = self.send_after {
if self.last_send.elapsed() >= sa {
return HeartbeatAction::SendEmpty;
}
}
HeartbeatAction::Idle
}
None => std::future::pending().await,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(start_paused = true)]
async fn sends_keepalive_when_idle() {
let mut hb = Heartbeat::new(Some(Duration::from_millis(100)), None);
assert_eq!(hb.tick().await, HeartbeatAction::SendEmpty);
}
#[tokio::test(start_paused = true)]
async fn detects_peer_timeout() {
let mut hb = Heartbeat::new(None, Some(Duration::from_millis(100)));
let mut action = hb.tick().await;
if action == HeartbeatAction::Idle {
action = hb.tick().await;
}
assert_eq!(action, HeartbeatAction::PeerTimedOut);
}
}