Skip to main content

scirs2_core/distributed/param_server/
fault_tolerance.rs

1//! Fault tolerance for the parameter server
2//!
3//! Provides heartbeat monitoring, worker failure detection, automatic
4//! checkpointing, and vector clocks for causal ordering.
5
6use crate::error::{CoreError, CoreResult, ErrorContext};
7
8use super::server::{ParameterServer, ServerCheckpoint};
9use super::types::ParamServerConfig;
10
11/// Configuration for the checkpointing subsystem
12#[derive(Debug, Clone)]
13pub struct CheckpointConfig {
14    /// How often to checkpoint (in steps); 0 disables auto-checkpointing
15    pub checkpoint_every: usize,
16    /// Maximum number of rolling checkpoints to retain
17    pub max_checkpoints: usize,
18}
19
20impl Default for CheckpointConfig {
21    fn default() -> Self {
22        Self {
23            checkpoint_every: 10,
24            max_checkpoints: 3,
25        }
26    }
27}
28
29/// Fault-tolerant PS with rolling checkpoints and heartbeat-based failure detection
30///
31/// This is a richer companion to [`FaultTolerantPS`] that keeps multiple rolling
32/// checkpoints and integrates a millisecond-granularity heartbeat timeout.
33#[derive(Debug)]
34pub struct FaultTolerantPs {
35    /// The underlying parameter server
36    server: ParameterServer,
37    /// Checkpoint configuration
38    checkpoint_config: CheckpointConfig,
39    /// Rolling checkpoint store (newest last)
40    checkpoints: std::collections::VecDeque<ServerCheckpoint>,
41    /// Per-worker heartbeat timestamp (as Instant)
42    worker_heartbeats: std::collections::HashMap<usize, std::time::Instant>,
43    /// Heartbeat timeout in milliseconds
44    heartbeat_timeout_ms: u64,
45}
46
47impl FaultTolerantPs {
48    /// Create a new `FaultTolerantPs` wrapping `server`
49    #[must_use]
50    pub fn new(server: ParameterServer, checkpoint_config: CheckpointConfig) -> Self {
51        let n = server.num_workers();
52        let mut worker_heartbeats = std::collections::HashMap::new();
53        for i in 0..n {
54            worker_heartbeats.insert(i, std::time::Instant::now());
55        }
56        Self {
57            server,
58            checkpoint_config,
59            checkpoints: std::collections::VecDeque::new(),
60            worker_heartbeats,
61            heartbeat_timeout_ms: 5000,
62        }
63    }
64
65    /// Get a reference to the underlying server
66    #[must_use]
67    pub fn server(&self) -> &ParameterServer {
68        &self.server
69    }
70
71    /// Get a mutable reference to the underlying server
72    pub fn server_mut(&mut self) -> &mut ParameterServer {
73        &mut self.server
74    }
75
76    /// Take a checkpoint, rolling off the oldest if `max_checkpoints` is exceeded
77    pub fn checkpoint(&mut self) -> CoreResult<()> {
78        let cp = self.server.checkpoint();
79        self.checkpoints.push_back(cp);
80        while self.checkpoints.len() > self.checkpoint_config.max_checkpoints {
81            self.checkpoints.pop_front();
82        }
83        Ok(())
84    }
85
86    /// Restore the parameter server from the most recent checkpoint
87    pub fn restore_latest(&mut self) -> CoreResult<()> {
88        let cp = self.checkpoints.back().ok_or_else(|| {
89            CoreError::ComputationError(ErrorContext::new(
90                "No checkpoint available to restore".to_string(),
91            ))
92        })?;
93        let config = self.server.config().clone();
94        self.server = ParameterServer::restore(cp, config)?;
95        Ok(())
96    }
97
98    /// Record a heartbeat for `worker_id` (updates timestamp to now)
99    pub fn heartbeat(&mut self, worker_id: usize) {
100        self.worker_heartbeats
101            .insert(worker_id, std::time::Instant::now());
102    }
103
104    /// Return the IDs of workers that have missed their heartbeat deadline
105    #[must_use]
106    pub fn detect_failed_workers(&self) -> Vec<usize> {
107        let timeout = std::time::Duration::from_millis(self.heartbeat_timeout_ms);
108        let now = std::time::Instant::now();
109        let mut failed = Vec::new();
110        for (worker_id, last_beat) in &self.worker_heartbeats {
111            if now.duration_since(*last_beat) > timeout {
112                failed.push(*worker_id);
113            }
114        }
115        failed.sort_unstable();
116        failed
117    }
118
119    /// Return the number of retained checkpoints
120    #[must_use]
121    pub fn n_checkpoints(&self) -> usize {
122        self.checkpoints.len()
123    }
124
125    /// Set heartbeat timeout in milliseconds
126    pub fn set_heartbeat_timeout_ms(&mut self, ms: u64) {
127        self.heartbeat_timeout_ms = ms;
128    }
129}
130
131/// Fault-tolerant wrapper around `ParameterServer`
132///
133/// Adds heartbeat monitoring, failure detection, automatic checkpointing,
134/// and recovery capabilities.
135#[derive(Debug)]
136pub struct FaultTolerantPS {
137    /// The underlying parameter server
138    server: ParameterServer,
139    /// Last checkpoint taken
140    last_checkpoint: Option<ServerCheckpoint>,
141    /// Step at which last checkpoint was taken
142    last_checkpoint_step: usize,
143}
144
145impl FaultTolerantPS {
146    /// Create a new fault-tolerant parameter server
147    #[must_use]
148    pub fn new(config: ParamServerConfig) -> Self {
149        Self {
150            server: ParameterServer::new(config),
151            last_checkpoint: None,
152            last_checkpoint_step: 0,
153        }
154    }
155
156    /// Get a reference to the underlying parameter server
157    #[must_use]
158    pub fn server(&self) -> &ParameterServer {
159        &self.server
160    }
161
162    /// Get a mutable reference to the underlying parameter server
163    pub fn server_mut(&mut self) -> &mut ParameterServer {
164        &mut self.server
165    }
166
167    /// Record a heartbeat from a worker
168    pub fn heartbeat(&mut self, worker_id: usize, timestamp: u64) -> CoreResult<()> {
169        let workers = self.server.workers_mut();
170        if worker_id >= workers.len() {
171            return Err(CoreError::ValueError(ErrorContext::new(format!(
172                "Unknown worker ID: {worker_id}"
173            ))));
174        }
175        workers[worker_id].last_heartbeat = timestamp;
176        workers[worker_id].is_alive = true;
177        Ok(())
178    }
179
180    /// Check for timed-out workers
181    ///
182    /// Returns the IDs of workers whose last heartbeat is older than
183    /// `current_time - timeout`.
184    #[must_use]
185    pub fn check_workers(&self, current_time: u64, timeout: u64) -> Vec<usize> {
186        self.server
187            .workers()
188            .iter()
189            .filter(|w| w.is_alive && current_time.saturating_sub(w.last_heartbeat) > timeout)
190            .map(|w| w.worker_id)
191            .collect()
192    }
193
194    /// Handle a worker failure by marking it as dead
195    ///
196    /// Any pending updates from this worker in the buffer are discarded
197    /// during the next barrier sync (BSP) since the worker count is reduced.
198    pub fn handle_worker_failure(&mut self, failed_worker: usize) -> CoreResult<()> {
199        let workers = self.server.workers_mut();
200        if failed_worker >= workers.len() {
201            return Err(CoreError::ValueError(ErrorContext::new(format!(
202                "Unknown worker ID: {failed_worker}"
203            ))));
204        }
205        if !workers[failed_worker].is_alive {
206            return Err(CoreError::ComputationError(ErrorContext::new(format!(
207                "Worker {failed_worker} is already marked as dead"
208            ))));
209        }
210        workers[failed_worker].is_alive = false;
211        Ok(())
212    }
213
214    /// Check if a checkpoint should be taken at this step, and if so, take it
215    ///
216    /// Returns the checkpoint if one was taken.
217    pub fn checkpoint_if_needed(&mut self, step: usize) -> Option<ServerCheckpoint> {
218        let interval = self.server.config().checkpoint_interval;
219        if interval == 0 {
220            return None;
221        }
222        if step > 0 && step % interval == 0 && step > self.last_checkpoint_step {
223            let cp = self.server.checkpoint();
224            self.last_checkpoint = Some(cp.clone());
225            self.last_checkpoint_step = step;
226            Some(cp)
227        } else {
228            None
229        }
230    }
231
232    /// Recover from a checkpoint
233    pub fn recover_from_checkpoint(
234        checkpoint: ServerCheckpoint,
235        config: ParamServerConfig,
236    ) -> CoreResult<Self> {
237        let server = ParameterServer::restore(&checkpoint, config)?;
238        Ok(Self {
239            server,
240            last_checkpoint: Some(checkpoint),
241            last_checkpoint_step: 0,
242        })
243    }
244
245    /// Get the last checkpoint if available
246    #[must_use]
247    pub fn last_checkpoint(&self) -> Option<&ServerCheckpoint> {
248        self.last_checkpoint.as_ref()
249    }
250}
251
252/// Vector clock for causal ordering of distributed updates
253///
254/// Each worker maintains a logical clock. The vector clock tracks all
255/// workers' clocks to determine causal relationships between events.
256#[derive(Debug, Clone, PartialEq, Eq)]
257pub struct VectorClock {
258    /// Logical clock value for each worker
259    clocks: Vec<u64>,
260}
261
262impl VectorClock {
263    /// Create a new vector clock for the given number of workers
264    #[must_use]
265    pub fn new(num_workers: usize) -> Self {
266        Self {
267            clocks: vec![0; num_workers],
268        }
269    }
270
271    /// Increment the clock for a specific worker (local event)
272    pub fn increment(&mut self, worker_id: usize) -> CoreResult<()> {
273        if worker_id >= self.clocks.len() {
274            return Err(CoreError::ValueError(ErrorContext::new(format!(
275                "Worker ID {worker_id} out of range (size = {})",
276                self.clocks.len()
277            ))));
278        }
279        self.clocks[worker_id] = self.clocks[worker_id].saturating_add(1);
280        Ok(())
281    }
282
283    /// Merge with another vector clock (receive event)
284    ///
285    /// Takes the element-wise maximum of both clocks.
286    pub fn merge(&mut self, other: &VectorClock) -> CoreResult<()> {
287        if self.clocks.len() != other.clocks.len() {
288            return Err(CoreError::DimensionError(ErrorContext::new(format!(
289                "Vector clock size mismatch: {} vs {}",
290                self.clocks.len(),
291                other.clocks.len()
292            ))));
293        }
294        for (mine, theirs) in self.clocks.iter_mut().zip(other.clocks.iter()) {
295            *mine = (*mine).max(*theirs);
296        }
297        Ok(())
298    }
299
300    /// Check if this clock causally happened before `other`
301    ///
302    /// Returns true if all components of self are <= the corresponding
303    /// components of other, and at least one is strictly less.
304    #[must_use]
305    pub fn happens_before(&self, other: &VectorClock) -> bool {
306        if self.clocks.len() != other.clocks.len() {
307            return false;
308        }
309        let mut all_leq = true;
310        let mut any_lt = false;
311        for (a, b) in self.clocks.iter().zip(other.clocks.iter()) {
312            if a > b {
313                all_leq = false;
314                break;
315            }
316            if a < b {
317                any_lt = true;
318            }
319        }
320        all_leq && any_lt
321    }
322
323    /// Check if two events are concurrent (neither happens-before the other)
324    #[must_use]
325    pub fn is_concurrent(&self, other: &VectorClock) -> bool {
326        !self.happens_before(other) && !other.happens_before(self) && self != other
327    }
328
329    /// Get the clock value for a specific worker
330    pub fn get(&self, worker_id: usize) -> CoreResult<u64> {
331        self.clocks.get(worker_id).copied().ok_or_else(|| {
332            CoreError::ValueError(ErrorContext::new(format!(
333                "Worker ID {worker_id} out of range"
334            )))
335        })
336    }
337
338    /// Get the number of workers tracked
339    #[must_use]
340    pub fn len(&self) -> usize {
341        self.clocks.len()
342    }
343
344    /// Check if the vector clock is empty
345    #[must_use]
346    pub fn is_empty(&self) -> bool {
347        self.clocks.is_empty()
348    }
349}
350
351#[cfg(test)]
352mod tests {
353    use super::*;
354    use crate::distributed::param_server::types::{ConsistencyModel, ParameterUpdate};
355
356    #[test]
357    fn test_heartbeat_and_check() {
358        let config = ParamServerConfig {
359            num_workers: 3,
360            ..ParamServerConfig::default()
361        };
362        let mut ftps = FaultTolerantPS::new(config);
363        ftps.server_mut().register_worker();
364        ftps.server_mut().register_worker();
365        ftps.server_mut().register_worker();
366
367        // All workers send heartbeats at t=10
368        for i in 0..3 {
369            ftps.heartbeat(i, 10).expect("heartbeat");
370        }
371
372        // At t=20 with timeout=15, no one is timed out
373        let timed_out = ftps.check_workers(20, 15);
374        assert!(timed_out.is_empty());
375
376        // At t=30 with timeout=15, all timed out (last heartbeat was at 10)
377        let timed_out = ftps.check_workers(30, 15);
378        assert_eq!(timed_out.len(), 3);
379    }
380
381    #[test]
382    fn test_handle_worker_failure() {
383        let config = ParamServerConfig {
384            num_workers: 2,
385            consistency: ConsistencyModel::ASP,
386            ..ParamServerConfig::default()
387        };
388        let mut ftps = FaultTolerantPS::new(config);
389        ftps.server_mut().register_worker();
390        ftps.server_mut().register_worker();
391
392        ftps.handle_worker_failure(0).expect("mark dead");
393
394        // Can't mark dead again
395        assert!(ftps.handle_worker_failure(0).is_err());
396
397        // Dead worker can't push
398        let result = ftps.server_mut().push(ParameterUpdate {
399            key: "k".into(),
400            values: vec![1.0],
401            worker_id: 0,
402            version: 1,
403        });
404        assert!(result.is_err());
405    }
406
407    #[test]
408    fn test_checkpoint_if_needed() {
409        let config = ParamServerConfig {
410            num_workers: 1,
411            checkpoint_interval: 5,
412            consistency: ConsistencyModel::ASP,
413            ..ParamServerConfig::default()
414        };
415        let mut ftps = FaultTolerantPS::new(config);
416        ftps.server_mut().register_worker();
417        ftps.server_mut().init_parameter("p", vec![1.0]);
418
419        // Step 3: no checkpoint
420        assert!(ftps.checkpoint_if_needed(3).is_none());
421        // Step 5: checkpoint
422        assert!(ftps.checkpoint_if_needed(5).is_some());
423        // Step 5 again: no double checkpoint
424        assert!(ftps.checkpoint_if_needed(5).is_none());
425        // Step 10: checkpoint
426        assert!(ftps.checkpoint_if_needed(10).is_some());
427    }
428
429    #[test]
430    fn test_recover_from_checkpoint() {
431        let config = ParamServerConfig {
432            num_workers: 1,
433            consistency: ConsistencyModel::ASP,
434            ..ParamServerConfig::default()
435        };
436        let mut ftps = FaultTolerantPS::new(config.clone());
437        let w0 = ftps.server_mut().register_worker();
438        ftps.server_mut().init_parameter("x", vec![99.0]);
439
440        let cp = ftps.server().checkpoint();
441        let recovered = FaultTolerantPS::recover_from_checkpoint(cp, config).expect("recover");
442        let (vals, _) = recovered.server().pull("x", w0).expect("pull");
443        assert!((vals[0] - 99.0).abs() < f64::EPSILON);
444    }
445
446    #[test]
447    fn test_vector_clock_basic() {
448        let mut vc1 = VectorClock::new(3);
449        vc1.increment(0).expect("inc");
450        vc1.increment(0).expect("inc");
451        assert_eq!(vc1.get(0).expect("get"), 2);
452        assert_eq!(vc1.get(1).expect("get"), 0);
453    }
454
455    #[test]
456    fn test_vector_clock_happens_before() {
457        let mut vc1 = VectorClock::new(2);
458        vc1.increment(0).expect("inc");
459
460        let mut vc2 = VectorClock::new(2);
461        vc2.increment(0).expect("inc");
462        vc2.increment(1).expect("inc");
463
464        // vc1 = [1, 0], vc2 = [1, 1] => vc1 happens-before vc2
465        assert!(vc1.happens_before(&vc2));
466        assert!(!vc2.happens_before(&vc1));
467    }
468
469    #[test]
470    fn test_vector_clock_concurrent() {
471        let mut vc1 = VectorClock::new(2);
472        vc1.increment(0).expect("inc");
473        // vc1 = [1, 0]
474
475        let mut vc2 = VectorClock::new(2);
476        vc2.increment(1).expect("inc");
477        // vc2 = [0, 1]
478
479        assert!(vc1.is_concurrent(&vc2));
480        assert!(vc2.is_concurrent(&vc1));
481    }
482
483    #[test]
484    fn test_vector_clock_merge() {
485        let mut vc1 = VectorClock::new(3);
486        vc1.increment(0).expect("inc");
487        vc1.increment(0).expect("inc");
488        // vc1 = [2, 0, 0]
489
490        let mut vc2 = VectorClock::new(3);
491        vc2.increment(1).expect("inc");
492        vc2.increment(2).expect("inc");
493        vc2.increment(2).expect("inc");
494        // vc2 = [0, 1, 2]
495
496        vc1.merge(&vc2).expect("merge");
497        // merged = [2, 1, 2]
498        assert_eq!(vc1.get(0).expect("get"), 2);
499        assert_eq!(vc1.get(1).expect("get"), 1);
500        assert_eq!(vc1.get(2).expect("get"), 2);
501    }
502
503    #[test]
504    fn test_vector_clock_size_mismatch() {
505        let vc1 = VectorClock::new(2);
506        let vc2 = VectorClock::new(3);
507        assert!(!vc1.happens_before(&vc2));
508    }
509
510    // ── FaultTolerantPs (rolling-checkpoint + Instant heartbeat) tests ──
511
512    fn make_ps_with_param() -> ParameterServer {
513        use super::super::types::{AggregationMethod, ConsistencyModel};
514        let config = ParamServerConfig {
515            num_workers: 2,
516            consistency: ConsistencyModel::ASP,
517            aggregation: AggregationMethod::Mean,
518            checkpoint_interval: 100,
519            replication_factor: 1,
520        };
521        let mut ps = ParameterServer::new(config);
522        ps.register_worker();
523        ps.register_worker();
524        ps.init_parameter("w", vec![1.0, 2.0, 3.0]);
525        ps
526    }
527
528    #[test]
529    fn test_fault_tolerant_ps_checkpoint_saved() {
530        let ps = make_ps_with_param();
531        let cc = CheckpointConfig {
532            checkpoint_every: 10,
533            max_checkpoints: 3,
534        };
535        let mut ftps = FaultTolerantPs::new(ps, cc);
536
537        assert_eq!(ftps.n_checkpoints(), 0);
538        ftps.checkpoint().expect("checkpoint");
539        assert_eq!(ftps.n_checkpoints(), 1);
540        ftps.checkpoint().expect("checkpoint 2");
541        assert_eq!(ftps.n_checkpoints(), 2);
542    }
543
544    #[test]
545    fn test_fault_tolerant_ps_rolling_eviction() {
546        let ps = make_ps_with_param();
547        let cc = CheckpointConfig {
548            checkpoint_every: 5,
549            max_checkpoints: 2,
550        };
551        let mut ftps = FaultTolerantPs::new(ps, cc);
552
553        ftps.checkpoint().expect("cp1");
554        ftps.checkpoint().expect("cp2");
555        ftps.checkpoint().expect("cp3");
556        // max_checkpoints = 2, so oldest is evicted
557        assert_eq!(ftps.n_checkpoints(), 2);
558    }
559
560    #[test]
561    fn test_fault_tolerant_ps_restore() {
562        let ps = make_ps_with_param();
563        let cc = CheckpointConfig::default();
564        let mut ftps = FaultTolerantPs::new(ps, cc);
565
566        ftps.checkpoint().expect("checkpoint before push");
567
568        // Modify parameter
569        use super::super::types::ParameterUpdate;
570        ftps.server_mut()
571            .push(ParameterUpdate {
572                key: "w".into(),
573                values: vec![99.0, 99.0, 99.0],
574                worker_id: 0,
575                version: 1,
576            })
577            .expect("push");
578
579        // Restore to checkpoint
580        ftps.restore_latest().expect("restore");
581
582        let (vals, _) = ftps.server().pull("w", 0).expect("pull after restore");
583        assert!((vals[0] - 1.0).abs() < f64::EPSILON);
584        assert!((vals[1] - 2.0).abs() < f64::EPSILON);
585    }
586
587    #[test]
588    fn test_fault_tolerant_ps_restore_no_checkpoint() {
589        let ps = make_ps_with_param();
590        let cc = CheckpointConfig::default();
591        let mut ftps = FaultTolerantPs::new(ps, cc);
592        // No checkpoint taken yet
593        assert!(ftps.restore_latest().is_err());
594    }
595
596    #[test]
597    fn test_fault_tolerant_ps_heartbeat_no_failure() {
598        let ps = make_ps_with_param();
599        let cc = CheckpointConfig::default();
600        let mut ftps = FaultTolerantPs::new(ps, cc);
601        // Set a very long timeout so no failures are detected
602        ftps.set_heartbeat_timeout_ms(60_000);
603        ftps.heartbeat(0);
604        ftps.heartbeat(1);
605        let failed = ftps.detect_failed_workers();
606        assert!(failed.is_empty(), "No workers should be failed: {failed:?}");
607    }
608
609    #[test]
610    fn test_fault_tolerant_ps_detect_failed_workers() {
611        let ps = make_ps_with_param();
612        let cc = CheckpointConfig::default();
613        let mut ftps = FaultTolerantPs::new(ps, cc);
614
615        // Set a zero-ms timeout so workers are immediately considered failed
616        ftps.set_heartbeat_timeout_ms(0);
617        // Brief sleep to ensure at least 1 ms has elapsed
618        std::thread::sleep(std::time::Duration::from_millis(2));
619
620        let failed = ftps.detect_failed_workers();
621        // Both workers should be detected as failed
622        assert!(!failed.is_empty(), "Workers should be detected as failed");
623    }
624
625    #[test]
626    fn test_checkpoint_config_default() {
627        let cc = CheckpointConfig::default();
628        assert_eq!(cc.checkpoint_every, 10);
629        assert_eq!(cc.max_checkpoints, 3);
630    }
631}