scirs2_core/distributed/param_server/
fault_tolerance.rs1use crate::error::{CoreError, CoreResult, ErrorContext};
7
8use super::server::{ParameterServer, ServerCheckpoint};
9use super::types::ParamServerConfig;
10
11#[derive(Debug, Clone)]
13pub struct CheckpointConfig {
14 pub checkpoint_every: usize,
16 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#[derive(Debug)]
34pub struct FaultTolerantPs {
35 server: ParameterServer,
37 checkpoint_config: CheckpointConfig,
39 checkpoints: std::collections::VecDeque<ServerCheckpoint>,
41 worker_heartbeats: std::collections::HashMap<usize, std::time::Instant>,
43 heartbeat_timeout_ms: u64,
45}
46
47impl FaultTolerantPs {
48 #[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 #[must_use]
67 pub fn server(&self) -> &ParameterServer {
68 &self.server
69 }
70
71 pub fn server_mut(&mut self) -> &mut ParameterServer {
73 &mut self.server
74 }
75
76 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 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 pub fn heartbeat(&mut self, worker_id: usize) {
100 self.worker_heartbeats
101 .insert(worker_id, std::time::Instant::now());
102 }
103
104 #[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 #[must_use]
121 pub fn n_checkpoints(&self) -> usize {
122 self.checkpoints.len()
123 }
124
125 pub fn set_heartbeat_timeout_ms(&mut self, ms: u64) {
127 self.heartbeat_timeout_ms = ms;
128 }
129}
130
131#[derive(Debug)]
136pub struct FaultTolerantPS {
137 server: ParameterServer,
139 last_checkpoint: Option<ServerCheckpoint>,
141 last_checkpoint_step: usize,
143}
144
145impl FaultTolerantPS {
146 #[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 #[must_use]
158 pub fn server(&self) -> &ParameterServer {
159 &self.server
160 }
161
162 pub fn server_mut(&mut self) -> &mut ParameterServer {
164 &mut self.server
165 }
166
167 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 #[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 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 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 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 #[must_use]
247 pub fn last_checkpoint(&self) -> Option<&ServerCheckpoint> {
248 self.last_checkpoint.as_ref()
249 }
250}
251
252#[derive(Debug, Clone, PartialEq, Eq)]
257pub struct VectorClock {
258 clocks: Vec<u64>,
260}
261
262impl VectorClock {
263 #[must_use]
265 pub fn new(num_workers: usize) -> Self {
266 Self {
267 clocks: vec![0; num_workers],
268 }
269 }
270
271 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 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 #[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 #[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 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 #[must_use]
340 pub fn len(&self) -> usize {
341 self.clocks.len()
342 }
343
344 #[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 for i in 0..3 {
369 ftps.heartbeat(i, 10).expect("heartbeat");
370 }
371
372 let timed_out = ftps.check_workers(20, 15);
374 assert!(timed_out.is_empty());
375
376 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 assert!(ftps.handle_worker_failure(0).is_err());
396
397 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 assert!(ftps.checkpoint_if_needed(3).is_none());
421 assert!(ftps.checkpoint_if_needed(5).is_some());
423 assert!(ftps.checkpoint_if_needed(5).is_none());
425 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 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 let mut vc2 = VectorClock::new(2);
476 vc2.increment(1).expect("inc");
477 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 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 vc1.merge(&vc2).expect("merge");
497 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 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 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 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 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 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 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 ftps.set_heartbeat_timeout_ms(0);
617 std::thread::sleep(std::time::Duration::from_millis(2));
619
620 let failed = ftps.detect_failed_workers();
621 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}