1use std::collections::HashMap;
13use std::sync::Arc;
14use tokio::sync::RwLock;
15
16#[derive(Debug, Clone, Copy)]
18pub struct HeartbeatConfig {
19 pub interval_ms: u64,
21 pub timeout_ms: u64,
23 pub max_missed: u32,
25}
26
27impl Default for HeartbeatConfig {
28 fn default() -> Self {
29 Self {
30 interval_ms: 30_000,
31 timeout_ms: 10_000,
32 max_missed: 3,
33 }
34 }
35}
36
37impl HeartbeatConfig {
38 pub fn new(interval_ms: u64, timeout_ms: u64, max_missed: u32) -> Self {
40 Self {
41 interval_ms,
42 timeout_ms,
43 max_missed,
44 }
45 }
46
47 pub fn validate(&self) -> Result<(), String> {
49 if self.interval_ms == 0 {
50 return Err("interval_ms must be > 0".to_string());
51 }
52 if self.timeout_ms == 0 {
53 return Err("timeout_ms must be > 0".to_string());
54 }
55 if self.timeout_ms >= self.interval_ms {
56 return Err("timeout_ms must be < interval_ms".to_string());
57 }
58 if self.max_missed == 0 {
59 return Err("max_missed must be > 0".to_string());
60 }
61 Ok(())
62 }
63}
64
65#[derive(Debug, Clone)]
67pub struct HeartbeatState {
68 pub connection_id: String,
70 pub last_ping_at: Option<i64>,
72 pub last_pong_at: Option<i64>,
74 pub missed_count: u32,
76 pub total_pings: u64,
78 pub total_pongs: u64,
80 pub is_dead: bool,
82}
83
84impl HeartbeatState {
85 pub fn new(connection_id: impl Into<String>) -> Self {
86 Self {
87 connection_id: connection_id.into(),
88 last_ping_at: None,
89 last_pong_at: None,
90 missed_count: 0,
91 total_pings: 0,
92 total_pongs: 0,
93 is_dead: false,
94 }
95 }
96
97 pub fn record_ping(&mut self, now_ms: i64) {
99 self.last_ping_at = Some(now_ms);
100 self.total_pings += 1;
101 }
102
103 pub fn record_pong(&mut self, now_ms: i64) -> bool {
105 self.last_pong_at = Some(now_ms);
106 self.total_pongs += 1;
107 let cleared = self.missed_count > 0;
108 self.missed_count = 0;
109 cleared
110 }
111
112 pub fn check_timeout(&mut self, now_ms: i64, config: &HeartbeatConfig) -> bool {
115 if self.is_dead {
116 return false;
117 }
118 let Some(last_ping) = self.last_ping_at else {
119 return false; };
121 if let Some(last_pong) = self.last_pong_at {
123 if last_pong >= last_ping {
124 return false;
125 }
126 }
127 if now_ms - last_ping < config.timeout_ms as i64 {
129 return false;
130 }
131 self.missed_count += 1;
133 if self.missed_count >= config.max_missed {
134 self.is_dead = true;
135 }
136 true
137 }
138
139 pub fn rtt_ms(&self) -> Option<i64> {
141 match (self.last_ping_at, self.last_pong_at) {
142 (Some(ping), Some(pong)) if pong >= ping => Some(pong - ping),
143 _ => None,
144 }
145 }
146
147 pub fn awaiting_pong(&self) -> bool {
149 match (self.last_ping_at, self.last_pong_at) {
150 (Some(ping), Some(pong)) => ping > pong,
151 (Some(_), None) => true,
152 _ => false,
153 }
154 }
155}
156
157#[derive(Debug)]
159pub struct HeartbeatTracker {
160 config: HeartbeatConfig,
161 states: Arc<RwLock<HashMap<String, HeartbeatState>>>,
162}
163
164impl HeartbeatTracker {
165 pub fn new(config: HeartbeatConfig) -> Self {
166 Self {
167 config,
168 states: Arc::new(RwLock::new(HashMap::new())),
169 }
170 }
171
172 pub fn config(&self) -> &HeartbeatConfig {
174 &self.config
175 }
176
177 pub async fn register(&self, connection_id: impl Into<String>) {
179 let id = connection_id.into();
180 let mut states = self.states.write().await;
181 states
182 .entry(id.clone())
183 .or_insert_with(|| HeartbeatState::new(id));
184 }
185
186 pub async fn register_new(&self, connection_id: impl Into<String>) {
188 let id = connection_id.into();
189 let mut states = self.states.write().await;
190 states.insert(id.clone(), HeartbeatState::new(id));
191 }
192
193 pub async fn unregister(&self, connection_id: &str) -> Option<HeartbeatState> {
195 let mut states = self.states.write().await;
196 states.remove(connection_id)
197 }
198
199 pub async fn record_ping(&self, connection_id: &str, now_ms: i64) -> bool {
201 let mut states = self.states.write().await;
202 if let Some(state) = states.get_mut(connection_id) {
203 state.record_ping(now_ms);
204 return true;
205 }
206 false
207 }
208
209 pub async fn record_pong(&self, connection_id: &str, now_ms: i64) -> bool {
211 let mut states = self.states.write().await;
212 if let Some(state) = states.get_mut(connection_id) {
213 state.record_pong(now_ms);
214 return true;
215 }
216 false
217 }
218
219 pub async fn check_timeouts(&self, now_ms: i64) -> (Vec<String>, Vec<String>) {
221 let mut states = self.states.write().await;
222 let mut newly_missed = Vec::new();
223 let mut newly_dead = Vec::new();
224 for (id, state) in states.iter_mut() {
225 let was_dead = state.is_dead;
226 let missed = state.check_timeout(now_ms, &self.config);
227 if missed {
229 newly_missed.push(id.clone());
230 }
231 if !was_dead && state.is_dead {
232 newly_dead.push(id.clone());
233 }
234 }
235 (newly_missed, newly_dead)
236 }
237
238 pub async fn state(&self, connection_id: &str) -> Option<HeartbeatState> {
240 let states = self.states.read().await;
241 states.get(connection_id).cloned()
242 }
243
244 pub async fn count(&self) -> usize {
246 let states = self.states.read().await;
247 states.len()
248 }
249
250 pub async fn dead_connections(&self) -> Vec<String> {
252 let states = self.states.read().await;
253 let mut dead: Vec<String> = states
254 .iter()
255 .filter(|(_, s)| s.is_dead)
256 .map(|(id, _)| id.clone())
257 .collect();
258 dead.sort();
259 dead
260 }
261
262 pub async fn purge_dead(&self) -> usize {
264 let mut states = self.states.write().await;
265 let before = states.len();
266 states.retain(|_, s| !s.is_dead);
267 before - states.len()
268 }
269
270 pub async fn rtts(&self) -> Vec<(String, i64)> {
272 let states = self.states.read().await;
273 let mut result: Vec<(String, i64)> = states
274 .iter()
275 .filter_map(|(id, s)| s.rtt_ms().map(|rtt| (id.clone(), rtt)))
276 .collect();
277 result.sort_by(|a, b| a.0.cmp(&b.0));
278 result
279 }
280}
281
282#[cfg(test)]
283mod tests {
284 use super::*;
285
286 #[test]
287 fn test_heartbeat_config_default() {
288 let cfg = HeartbeatConfig::default();
289 assert_eq!(cfg.interval_ms, 30_000);
290 assert_eq!(cfg.timeout_ms, 10_000);
291 assert_eq!(cfg.max_missed, 3);
292 }
293
294 #[test]
295 fn test_heartbeat_config_validate_ok() {
296 let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
297 assert!(cfg.validate().is_ok());
298 assert_eq!(cfg.interval_ms, 30_000, "validate 不应修改 interval_ms");
299 assert_eq!(cfg.timeout_ms, 10_000, "validate 不应修改 timeout_ms");
300 assert_eq!(cfg.max_missed, 3, "validate 不应修改 max_missed");
301 }
302
303 #[test]
304 fn test_heartbeat_config_validate_zero_interval() {
305 let cfg = HeartbeatConfig::new(0, 10_000, 3);
306 assert!(cfg.validate().is_err());
307 }
308
309 #[test]
310 fn test_heartbeat_config_validate_zero_timeout() {
311 let cfg = HeartbeatConfig::new(30_000, 0, 3);
312 assert!(cfg.validate().is_err());
313 }
314
315 #[test]
316 fn test_heartbeat_config_validate_timeout_ge_interval() {
317 let cfg = HeartbeatConfig::new(10_000, 10_000, 3);
318 assert!(cfg.validate().is_err());
319 let cfg2 = HeartbeatConfig::new(10_000, 20_000, 3);
320 assert!(cfg2.validate().is_err());
321 }
322
323 #[test]
324 fn test_heartbeat_config_validate_zero_max_missed() {
325 let cfg = HeartbeatConfig::new(30_000, 10_000, 0);
326 assert!(cfg.validate().is_err());
327 }
328
329 #[test]
330 fn test_heartbeat_state_new_defaults() {
331 let state = HeartbeatState::new("c1");
332 assert_eq!(state.connection_id, "c1");
333 assert!(state.last_ping_at.is_none());
334 assert!(state.last_pong_at.is_none());
335 assert_eq!(state.missed_count, 0);
336 assert_eq!(state.total_pings, 0);
337 assert_eq!(state.total_pongs, 0);
338 assert!(!state.is_dead);
339 }
340
341 #[test]
342 fn test_record_ping_updates_state() {
343 let mut state = HeartbeatState::new("c1");
344 state.record_ping(1000);
345 assert_eq!(state.last_ping_at, Some(1000));
346 assert_eq!(state.total_pings, 1);
347 assert!(state.awaiting_pong());
348 }
349
350 #[test]
351 fn test_record_pong_clears_missed_count() {
352 let mut state = HeartbeatState::new("c1");
353 state.record_ping(1000);
354 state.missed_count = 2;
355 let cleared = state.record_pong(2000);
356 assert!(cleared);
357 assert_eq!(state.missed_count, 0);
358 assert_eq!(state.total_pongs, 1);
359 assert!(!state.awaiting_pong());
360 }
361
362 #[test]
363 fn test_record_pong_no_missed_returns_false() {
364 let mut state = HeartbeatState::new("c1");
365 state.record_ping(1000);
366 let cleared = state.record_pong(2000);
367 assert!(!cleared); }
369
370 #[test]
371 fn test_rtt_ms_calculated_correctly() {
372 let mut state = HeartbeatState::new("c1");
373 state.record_ping(1000);
374 state.record_pong(1500);
375 assert_eq!(state.rtt_ms(), Some(500));
376 }
377
378 #[test]
379 fn test_rtt_ms_none_without_pong() {
380 let mut state = HeartbeatState::new("c1");
381 state.record_ping(1000);
382 assert_eq!(state.rtt_ms(), None);
383 }
384
385 #[test]
386 fn test_rtt_ms_none_without_ping() {
387 let state = HeartbeatState::new("c1");
388 assert_eq!(state.rtt_ms(), None);
389 }
390
391 #[test]
392 fn test_awaiting_pong_states() {
393 let mut state = HeartbeatState::new("c1");
394 assert!(!state.awaiting_pong());
395 state.record_ping(1000);
396 assert!(state.awaiting_pong());
397 state.record_pong(2000);
398 assert!(!state.awaiting_pong());
399 }
400
401 #[test]
402 fn test_check_timeout_no_ping_returns_false() {
403 let mut state = HeartbeatState::new("c1");
404 let cfg = HeartbeatConfig::default();
405 assert!(!state.check_timeout(100_000, &cfg));
406 }
407
408 #[test]
409 fn test_check_timeout_within_window_returns_false() {
410 let mut state = HeartbeatState::new("c1");
411 let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
412 state.record_ping(1000);
413 assert!(!state.check_timeout(6_000, &cfg));
415 assert_eq!(state.missed_count, 0);
416 }
417
418 #[test]
419 fn test_check_timeout_expired_increments_missed() {
420 let mut state = HeartbeatState::new("c1");
421 let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
422 state.record_ping(1000);
423 assert!(state.check_timeout(16_000, &cfg));
425 assert_eq!(state.missed_count, 1);
426 assert!(!state.is_dead);
427 }
428
429 #[test]
430 fn test_check_timeout_marks_dead_after_max_missed() {
431 let mut state = HeartbeatState::new("c1");
432 let cfg = HeartbeatConfig::new(30_000, 10_000, 2);
433 state.record_ping(1000);
434 state.check_timeout(16_000, &cfg); assert!(!state.is_dead);
436 state.record_ping(40_000);
438 state.check_timeout(56_000, &cfg); assert!(state.is_dead);
440 }
441
442 #[test]
443 fn test_check_timeout_dead_state_returns_false() {
444 let mut state = HeartbeatState::new("c1");
445 let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
446 state.record_ping(1000);
447 state.check_timeout(16_000, &cfg);
448 assert!(state.is_dead);
449 assert!(!state.check_timeout(100_000, &cfg));
451 }
452
453 #[tokio::test]
454 async fn test_tracker_register_new() {
455 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
456 tracker.register_new("c1").await;
457 assert_eq!(tracker.count().await, 1);
458 let state = tracker.state("c1").await.unwrap();
459 assert_eq!(state.connection_id, "c1");
460 }
461
462 #[tokio::test]
463 async fn test_tracker_unregister() {
464 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
465 tracker.register_new("c1").await;
466 let removed = tracker.unregister("c1").await;
467 assert!(removed.is_some());
468 assert_eq!(tracker.count().await, 0);
469 }
470
471 #[tokio::test]
472 async fn test_tracker_record_ping_and_pong() {
473 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
474 tracker.register_new("c1").await;
475 assert!(tracker.record_ping("c1", 1000).await);
476 assert!(tracker.record_pong("c1", 1500).await);
477 let state = tracker.state("c1").await.unwrap();
478 assert_eq!(state.total_pings, 1);
479 assert_eq!(state.total_pongs, 1);
480 assert_eq!(state.rtt_ms(), Some(500));
481 }
482
483 #[tokio::test]
484 async fn test_tracker_record_ping_unknown_returns_false() {
485 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
486 assert!(!tracker.record_ping("ghost", 1000).await);
487 }
488
489 #[tokio::test]
490 async fn test_tracker_record_pong_unknown_returns_false() {
491 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
492 assert!(!tracker.record_pong("ghost", 1000).await);
493 }
494
495 #[tokio::test]
496 async fn test_tracker_check_timeouts_detects_missed() {
497 let cfg = HeartbeatConfig::new(30_000, 10_000, 3);
498 let tracker = HeartbeatTracker::new(cfg);
499 tracker.register_new("c1").await;
500 tracker.register_new("c2").await;
501 tracker.record_ping("c1", 1000).await;
502 tracker.record_ping("c2", 1000).await;
503 tracker.record_pong("c2", 1500).await;
505 let (missed, dead) = tracker.check_timeouts(16_000).await;
507 assert_eq!(missed.len(), 1);
508 assert_eq!(missed[0], "c1");
509 assert!(dead.is_empty());
510 }
511
512 #[tokio::test]
513 async fn test_tracker_check_timeouts_detects_dead() {
514 let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
515 let tracker = HeartbeatTracker::new(cfg);
516 tracker.register_new("c1").await;
517 tracker.record_ping("c1", 1000).await;
518 let (missed, dead) = tracker.check_timeouts(16_000).await;
519 assert_eq!(missed.len(), 1);
520 assert_eq!(dead.len(), 1);
521 assert_eq!(dead[0], "c1");
522 }
523
524 #[tokio::test]
525 async fn test_tracker_dead_connections() {
526 let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
527 let tracker = HeartbeatTracker::new(cfg);
528 tracker.register_new("c1").await;
529 tracker.register_new("c2").await;
530 tracker.record_ping("c1", 1000).await;
531 tracker.check_timeouts(16_000).await; let dead = tracker.dead_connections().await;
533 assert_eq!(dead, vec!["c1"]);
534 }
535
536 #[tokio::test]
537 async fn test_tracker_purge_dead() {
538 let cfg = HeartbeatConfig::new(30_000, 10_000, 1);
539 let tracker = HeartbeatTracker::new(cfg);
540 tracker.register_new("c1").await;
541 tracker.register_new("c2").await;
542 tracker.record_ping("c1", 1000).await;
543 tracker.check_timeouts(16_000).await;
544 let purged = tracker.purge_dead().await;
545 assert_eq!(purged, 1);
546 assert_eq!(tracker.count().await, 1);
547 }
548
549 #[tokio::test]
550 async fn test_tracker_rtts() {
551 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
552 tracker.register_new("c1").await;
553 tracker.register_new("c2").await;
554 tracker.record_ping("c1", 1000).await;
555 tracker.record_pong("c1", 1500).await;
556 tracker.record_ping("c2", 2000).await;
557 tracker.record_pong("c2", 2800).await;
558 let rtts = tracker.rtts().await;
559 assert_eq!(rtts.len(), 2);
560 assert_eq!(rtts[0], ("c1".to_string(), 500));
561 assert_eq!(rtts[1], ("c2".to_string(), 800));
562 }
563
564 #[tokio::test]
565 async fn test_tracker_rtts_excludes_no_rtt() {
566 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
567 tracker.register_new("c1").await;
568 tracker.register_new("c2").await;
569 tracker.record_ping("c1", 1000).await;
570 tracker.record_pong("c1", 1500).await;
571 tracker.record_ping("c2", 2000).await;
573 let rtts = tracker.rtts().await;
574 assert_eq!(rtts.len(), 1);
575 assert_eq!(rtts[0].0, "c1");
576 }
577
578 #[tokio::test]
579 async fn test_tracker_multiple_pings_accumulate_stats() {
580 let tracker = HeartbeatTracker::new(HeartbeatConfig::default());
581 tracker.register_new("c1").await;
582 tracker.record_ping("c1", 1000).await;
583 tracker.record_pong("c1", 1500).await;
584 tracker.record_ping("c1", 2000).await;
585 tracker.record_pong("c1", 2200).await;
586 let state = tracker.state("c1").await.unwrap();
587 assert_eq!(state.total_pings, 2);
588 assert_eq!(state.total_pongs, 2);
589 assert_eq!(state.rtt_ms(), Some(200)); }
591}