Skip to main content

sz_orm_websocket/
heartbeat.rs

1//! # 心跳机制(Ping/Pong)
2//!
3//! 实现 WebSocket 心跳保活:周期性发送 Ping,等待 Pong 应答;
4//! 超时未应答则判定连接失活并触发断开。
5//!
6//! ## 主要类型
7//!
8//! - [`HeartbeatConfig`] — 心跳配置
9//! - [`HeartbeatState`] — 单连接的心跳状态
10//! - [`HeartbeatTracker`] — 多连接心跳跟踪器
11
12use std::collections::HashMap;
13use std::sync::Arc;
14use tokio::sync::RwLock;
15
16/// 心跳配置
17#[derive(Debug, Clone, Copy)]
18pub struct HeartbeatConfig {
19    /// Ping 发送间隔(毫秒)
20    pub interval_ms: u64,
21    /// 等待 Pong 的超时时间(毫秒)
22    pub timeout_ms: u64,
23    /// 连续未收到 Pong 的最大次数,超过则判定连接失活
24    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    /// 创建自定义配置
39    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    /// 校验配置合法性
48    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/// 单连接的心跳状态
66#[derive(Debug, Clone)]
67pub struct HeartbeatState {
68    /// 连接 ID
69    pub connection_id: String,
70    /// 上次发送 Ping 的时间戳(毫秒),None 表示尚未发送
71    pub last_ping_at: Option<i64>,
72    /// 上次收到 Pong 的时间戳(毫秒)
73    pub last_pong_at: Option<i64>,
74    /// 连续未收到 Pong 的次数
75    pub missed_count: u32,
76    /// 总共发送的 Ping 次数
77    pub total_pings: u64,
78    /// 总共收到的 Pong 次数
79    pub total_pongs: u64,
80    /// 是否被判定为失活
81    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    /// 记录发送了一次 Ping
98    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    /// 记录收到了一次 Pong。返回是否清除了未应答计数。
104    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    /// 检查是否超时未收到 Pong。
113    /// 返回 true 表示本次检查新增了一次未应答。
114    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; // 尚未发送 Ping
120        };
121        // 若已收到 Pong 且 Pong 时间 >= Ping 时间,说明本次 Ping 已被应答,不超时
122        if let Some(last_pong) = self.last_pong_at {
123            if last_pong >= last_ping {
124                return false;
125            }
126        }
127        // 未超时
128        if now_ms - last_ping < config.timeout_ms as i64 {
129            return false;
130        }
131        // 超时:增加 missed_count
132        self.missed_count += 1;
133        if self.missed_count >= config.max_missed {
134            self.is_dead = true;
135        }
136        true
137    }
138
139    /// 当前 RTT 估计(毫秒)。需要 last_ping 和 last_pong 都存在。
140    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    /// 是否正在等待 Pong 应答
148    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/// 多连接心跳跟踪器
158#[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    /// 获取配置
173    pub fn config(&self) -> &HeartbeatConfig {
174        &self.config
175    }
176
177    /// 注册一个连接(若已存在则保留原状态)
178    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    /// 注册一个连接(确保新建状态)
187    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    /// 注销连接
194    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    /// 记录发送 Ping
200    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    /// 记录收到 Pong
210    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    /// 检查所有连接的超时状态,返回 (本次新增超时的连接列表, 本次被判定为失活的连接列表)
220    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            // 新增超时的连接(含本次同时被判定为失活的)
228            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    /// 获取指定连接的心跳状态
239    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    /// 当前跟踪的连接数
245    pub async fn count(&self) -> usize {
246        let states = self.states.read().await;
247        states.len()
248    }
249
250    /// 获取所有失活连接的 ID
251    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    /// 清理所有失活连接,返回被清理的数量
263    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    /// 获取所有连接的 RTT(毫秒),仅返回有有效 RTT 的
271    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); // missed_count 本来就是 0
368    }
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        // 5 秒后检查(未超时)
414        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        // 15 秒后检查(已超时)
424        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); // missed=1
435        assert!(!state.is_dead);
436        // 再次检查(模拟下一个周期)
437        state.record_ping(40_000);
438        state.check_timeout(56_000, &cfg); // missed=2
439        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        // 已死,再次检查应返回 false
450        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        // c2 立即回复 Pong
504        tracker.record_pong("c2", 1500).await;
505        // 15 秒后检查:c1 超时,c2 正常
506        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; // c1 失活
532        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        // c2 仅发送 Ping 未收到 Pong
572        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)); // 最近一次 RTT
590    }
591}