use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone, Copy)]
pub struct HeartbeatConfig {
pub interval_ms: u64,
pub timeout_ms: u64,
pub max_missed: u32,
}
impl Default for HeartbeatConfig {
fn default() -> Self {
Self {
interval_ms: 30_000,
timeout_ms: 10_000,
max_missed: 3,
}
}
}
impl HeartbeatConfig {
pub fn new(interval_ms: u64, timeout_ms: u64, max_missed: u32) -> Self {
Self {
interval_ms,
timeout_ms,
max_missed,
}
}
pub fn validate(&self) -> Result<(), String> {
if self.interval_ms == 0 {
return Err("interval_ms must be > 0".to_string());
}
if self.timeout_ms == 0 {
return Err("timeout_ms must be > 0".to_string());
}
if self.timeout_ms >= self.interval_ms {
return Err("timeout_ms must be < interval_ms".to_string());
}
if self.max_missed == 0 {
return Err("max_missed must be > 0".to_string());
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct HeartbeatState {
pub connection_id: String,
pub last_ping_at: Option<i64>,
pub last_pong_at: Option<i64>,
pub missed_count: u32,
pub total_pings: u64,
pub total_pongs: u64,
pub is_dead: bool,
}
impl HeartbeatState {
pub fn new(connection_id: impl Into<String>) -> Self {
Self {
connection_id: connection_id.into(),
last_ping_at: None,
last_pong_at: None,
missed_count: 0,
total_pings: 0,
total_pongs: 0,
is_dead: false,
}
}
pub fn record_ping(&mut self, now_ms: i64) {
self.last_ping_at = Some(now_ms);
self.total_pings += 1;
}
pub fn record_pong(&mut self, now_ms: i64) -> bool {
self.last_pong_at = Some(now_ms);
self.total_pongs += 1;
let cleared = self.missed_count > 0;
self.missed_count = 0;
cleared
}
pub fn check_timeout(&mut self, now_ms: i64, config: &HeartbeatConfig) -> bool {
if self.is_dead {
return false;
}
let Some(last_ping) = self.last_ping_at else {
return false; };
if let Some(last_pong) = self.last_pong_at {
if last_pong >= last_ping {
return false;
}
}
if now_ms - last_ping < config.timeout_ms as i64 {
return false;
}
self.missed_count += 1;
if self.missed_count >= config.max_missed {
self.is_dead = true;
}
true
}
pub fn rtt_ms(&self) -> Option<i64> {
match (self.last_ping_at, self.last_pong_at) {
(Some(ping), Some(pong)) if pong >= ping => Some(pong - ping),
_ => None,
}
}
pub fn awaiting_pong(&self) -> bool {
match (self.last_ping_at, self.last_pong_at) {
(Some(ping), Some(pong)) => ping > pong,
(Some(_), None) => true,
_ => false,
}
}
}
#[derive(Debug)]
pub struct HeartbeatTracker {
config: HeartbeatConfig,
states: Arc<RwLock<HashMap<String, HeartbeatState>>>,
}
impl HeartbeatTracker {
pub fn new(config: HeartbeatConfig) -> Self {
Self {
config,
states: Arc::new(RwLock::new(HashMap::new())),
}
}
pub fn config(&self) -> &HeartbeatConfig {
&self.config
}
pub async fn register(&self, connection_id: impl Into<String>) {
let id = connection_id.into();
let mut states = self.states.write().await;
states
.entry(id.clone())
.or_insert_with(|| HeartbeatState::new(id));
}
pub async fn register_new(&self, connection_id: impl Into<String>) {
let id = connection_id.into();
let mut states = self.states.write().await;
states.insert(id.clone(), HeartbeatState::new(id));
}
pub async fn unregister(&self, connection_id: &str) -> Option<HeartbeatState> {
let mut states = self.states.write().await;
states.remove(connection_id)
}
pub async fn record_ping(&self, connection_id: &str, now_ms: i64) -> bool {
let mut states = self.states.write().await;
if let Some(state) = states.get_mut(connection_id) {
state.record_ping(now_ms);
return true;
}
false
}
pub async fn record_pong(&self, connection_id: &str, now_ms: i64) -> bool {
let mut states = self.states.write().await;
if let Some(state) = states.get_mut(connection_id) {
state.record_pong(now_ms);
return true;
}
false
}
pub async fn check_timeouts(&self, now_ms: i64) -> (Vec<String>, Vec<String>) {
let mut states = self.states.write().await;
let mut newly_missed = Vec::new();
let mut newly_dead = Vec::new();
for (id, state) in states.iter_mut() {
let was_dead = state.is_dead;
let missed = state.check_timeout(now_ms, &self.config);
if missed {
newly_missed.push(id.clone());
}
if !was_dead && state.is_dead {
newly_dead.push(id.clone());
}
}
(newly_missed, newly_dead)
}
pub async fn state(&self, connection_id: &str) -> Option<HeartbeatState> {
let states = self.states.read().await;
states.get(connection_id).cloned()
}
pub async fn count(&self) -> usize {
let states = self.states.read().await;
states.len()
}
pub async fn dead_connections(&self) -> Vec<String> {
let states = self.states.read().await;
let mut dead: Vec<String> = states
.iter()
.filter(|(_, s)| s.is_dead)
.map(|(id, _)| id.clone())
.collect();
dead.sort();
dead
}
pub async fn purge_dead(&self) -> usize {
let mut states = self.states.write().await;
let before = states.len();
states.retain(|_, s| !s.is_dead);
before - states.len()
}
pub async fn rtts(&self) -> Vec<(String, i64)> {
let states = self.states.read().await;
let mut result: Vec<(String, i64)> = states
.iter()
.filter_map(|(id, s)| s.rtt_ms().map(|rtt| (id.clone(), rtt)))
.collect();
result.sort_by(|a, b| a.0.cmp(&b.0));
result
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_heartbeat_config_default() {
let cfg = HeartbeatConfig::default();
assert_eq!(cfg.interval_ms, 30_000);
assert_eq!(cfg.timeout_ms, 10_000);
assert_eq!(cfg.max_missed, 3);
}
#[test]
fn test_heartbeat_config_validate_ok() {
let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
assert!(cfg.validate().is_ok());
assert_eq!(cfg.interval_ms, 30_000, "validate 不应修改 interval_ms");
assert_eq!(cfg.timeout_ms, 10_000, "validate 不应修改 timeout_ms");
assert_eq!(cfg.max_missed, 3, "validate 不应修改 max_missed");
}
#[test]
fn test_heartbeat_config_validate_zero_interval() {
let cfg = HeartbeatConfig::new(0, 10_000, 3);
assert!(cfg.validate().is_err());
}
#[test]
fn test_heartbeat_config_validate_zero_timeout() {
let cfg = HeartbeatConfig::new(30_000, 0, 3);
assert!(cfg.validate().is_err());
}
#[test]
fn test_heartbeat_config_validate_timeout_ge_interval() {
let cfg = HeartbeatConfig::new(10_000, 10_000, 3);
assert!(cfg.validate().is_err());
let cfg2 = HeartbeatConfig::new(10_000, 20_000, 3);
assert!(cfg2.validate().is_err());
}
#[test]
fn test_heartbeat_config_validate_zero_max_missed() {
let cfg = HeartbeatConfig::new(30_000, 10_000, 0);
assert!(cfg.validate().is_err());
}
#[test]
fn test_heartbeat_state_new_defaults() {
let state = HeartbeatState::new("c1");
assert_eq!(state.connection_id, "c1");
assert!(state.last_ping_at.is_none());
assert!(state.last_pong_at.is_none());
assert_eq!(state.missed_count, 0);
assert_eq!(state.total_pings, 0);
assert_eq!(state.total_pongs, 0);
assert!(!state.is_dead);
}
#[test]
fn test_record_ping_updates_state() {
let mut state = HeartbeatState::new("c1");
state.record_ping(1000);
assert_eq!(state.last_ping_at, Some(1000));
assert_eq!(state.total_pings, 1);
assert!(state.awaiting_pong());
}
#[test]
fn test_record_pong_clears_missed_count() {
let mut state = HeartbeatState::new("c1");
state.record_ping(1000);
state.missed_count = 2;
let cleared = state.record_pong(2000);
assert!(cleared);
assert_eq!(state.missed_count, 0);
assert_eq!(state.total_pongs, 1);
assert!(!state.awaiting_pong());
}
#[test]
fn test_record_pong_no_missed_returns_false() {
let mut state = HeartbeatState::new("c1");
state.record_ping(1000);
let cleared = state.record_pong(2000);
assert!(!cleared); }
#[test]
fn test_rtt_ms_calculated_correctly() {
let mut state = HeartbeatState::new("c1");
state.record_ping(1000);
state.record_pong(1500);
assert_eq!(state.rtt_ms(), Some(500));
}
#[test]
fn test_rtt_ms_none_without_pong() {
let mut state = HeartbeatState::new("c1");
state.record_ping(1000);
assert_eq!(state.rtt_ms(), None);
}
#[test]
fn test_rtt_ms_none_without_ping() {
let state = HeartbeatState::new("c1");
assert_eq!(state.rtt_ms(), None);
}
#[test]
fn test_awaiting_pong_states() {
let mut state = HeartbeatState::new("c1");
assert!(!state.awaiting_pong());
state.record_ping(1000);
assert!(state.awaiting_pong());
state.record_pong(2000);
assert!(!state.awaiting_pong());
}
#[test]
fn test_check_timeout_no_ping_returns_false() {
let mut state = HeartbeatState::new("c1");
let cfg = HeartbeatConfig::default();
assert!(!state.check_timeout(100_000, &cfg));
}
#[test]
fn test_check_timeout_within_window_returns_false() {
let mut state = HeartbeatState::new("c1");
let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
state.record_ping(1000);
assert!(!state.check_timeout(6_000, &cfg));
assert_eq!(state.missed_count, 0);
}
#[test]
fn test_check_timeout_expired_increments_missed() {
let mut state = HeartbeatState::new("c1");
let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
state.record_ping(1000);
assert!(state.check_timeout(16_000, &cfg));
assert_eq!(state.missed_count, 1);
assert!(!state.is_dead);
}
#[test]
fn test_check_timeout_marks_dead_after_max_missed() {
let mut state = HeartbeatState::new("c1");
let cfg = HeartbeatConfig::new(30_000, 10_000, 2);
state.record_ping(1000);
state.check_timeout(16_000, &cfg); assert!(!state.is_dead);
state.record_ping(40_000);
state.check_timeout(56_000, &cfg); assert!(state.is_dead);
}
#[test]
fn test_check_timeout_dead_state_returns_false() {
let mut state = HeartbeatState::new("c1");
let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
state.record_ping(1000);
state.check_timeout(16_000, &cfg);
assert!(state.is_dead);
assert!(!state.check_timeout(100_000, &cfg));
}
#[tokio::test]
async fn test_tracker_register_new() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
tracker.register_new("c1").await;
assert_eq!(tracker.count().await, 1);
let state = tracker.state("c1").await.unwrap();
assert_eq!(state.connection_id, "c1");
}
#[tokio::test]
async fn test_tracker_unregister() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
tracker.register_new("c1").await;
let removed = tracker.unregister("c1").await;
assert!(removed.is_some());
assert_eq!(tracker.count().await, 0);
}
#[tokio::test]
async fn test_tracker_record_ping_and_pong() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
tracker.register_new("c1").await;
assert!(tracker.record_ping("c1", 1000).await);
assert!(tracker.record_pong("c1", 1500).await);
let state = tracker.state("c1").await.unwrap();
assert_eq!(state.total_pings, 1);
assert_eq!(state.total_pongs, 1);
assert_eq!(state.rtt_ms(), Some(500));
}
#[tokio::test]
async fn test_tracker_record_ping_unknown_returns_false() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
assert!(!tracker.record_ping("ghost", 1000).await);
}
#[tokio::test]
async fn test_tracker_record_pong_unknown_returns_false() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
assert!(!tracker.record_pong("ghost", 1000).await);
}
#[tokio::test]
async fn test_tracker_check_timeouts_detects_missed() {
let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
let tracker = HeartbeatTracker::new(cfg);
tracker.register_new("c1").await;
tracker.register_new("c2").await;
tracker.record_ping("c1", 1000).await;
tracker.record_ping("c2", 1000).await;
tracker.record_pong("c2", 1500).await;
let (missed, dead) = tracker.check_timeouts(16_000).await;
assert_eq!(missed.len(), 1);
assert_eq!(missed[0], "c1");
assert!(dead.is_empty());
}
#[tokio::test]
async fn test_tracker_check_timeouts_detects_dead() {
let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
let tracker = HeartbeatTracker::new(cfg);
tracker.register_new("c1").await;
tracker.record_ping("c1", 1000).await;
let (missed, dead) = tracker.check_timeouts(16_000).await;
assert_eq!(missed.len(), 1);
assert_eq!(dead.len(), 1);
assert_eq!(dead[0], "c1");
}
#[tokio::test]
async fn test_tracker_dead_connections() {
let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
let tracker = HeartbeatTracker::new(cfg);
tracker.register_new("c1").await;
tracker.register_new("c2").await;
tracker.record_ping("c1", 1000).await;
tracker.check_timeouts(16_000).await; let dead = tracker.dead_connections().await;
assert_eq!(dead, vec!["c1"]);
}
#[tokio::test]
async fn test_tracker_purge_dead() {
let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
let tracker = HeartbeatTracker::new(cfg);
tracker.register_new("c1").await;
tracker.register_new("c2").await;
tracker.record_ping("c1", 1000).await;
tracker.check_timeouts(16_000).await;
let purged = tracker.purge_dead().await;
assert_eq!(purged, 1);
assert_eq!(tracker.count().await, 1);
}
#[tokio::test]
async fn test_tracker_rtts() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
tracker.register_new("c1").await;
tracker.register_new("c2").await;
tracker.record_ping("c1", 1000).await;
tracker.record_pong("c1", 1500).await;
tracker.record_ping("c2", 2000).await;
tracker.record_pong("c2", 2800).await;
let rtts = tracker.rtts().await;
assert_eq!(rtts.len(), 2);
assert_eq!(rtts[0], ("c1".to_string(), 500));
assert_eq!(rtts[1], ("c2".to_string(), 800));
}
#[tokio::test]
async fn test_tracker_rtts_excludes_no_rtt() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
tracker.register_new("c1").await;
tracker.register_new("c2").await;
tracker.record_ping("c1", 1000).await;
tracker.record_pong("c1", 1500).await;
tracker.record_ping("c2", 2000).await;
let rtts = tracker.rtts().await;
assert_eq!(rtts.len(), 1);
assert_eq!(rtts[0].0, "c1");
}
#[tokio::test]
async fn test_tracker_multiple_pings_accumulate_stats() {
let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
tracker.register_new("c1").await;
tracker.record_ping("c1", 1000).await;
tracker.record_pong("c1", 1500).await;
tracker.record_ping("c1", 2000).await;
tracker.record_pong("c1", 2200).await;
let state = tracker.state("c1").await.unwrap();
assert_eq!(state.total_pings, 2);
assert_eq!(state.total_pongs, 2);
assert_eq!(state.rtt_ms(), Some(200)); }
}