use diskann::ANNResult;
use tracing::info;
use super::{Progress, WorkStage};
pub trait CheckpointManager: Send + Sync + CheckpointManagerClone {
fn get_resumption_point(&self, stage: WorkStage) -> ANNResult<Option<usize>>;
fn update(&mut self, progress: Progress, next_stage: WorkStage) -> ANNResult<()>;
fn mark_as_invalid(&mut self) -> ANNResult<()>;
}
pub trait CheckpointManagerExt {
fn execute_stage<F, S, U>(
&mut self,
stage: WorkStage,
next_stage: WorkStage,
operation: F,
skip_handler: S,
) -> ANNResult<U>
where
F: FnOnce() -> ANNResult<U>,
S: FnOnce() -> ANNResult<U>;
}
impl<T: ?Sized> CheckpointManagerExt for T
where
T: CheckpointManager,
{
fn execute_stage<F, S, U>(
&mut self,
stage: WorkStage,
next_stage: WorkStage,
operation: F,
skip_handler: S,
) -> ANNResult<U>
where
F: FnOnce() -> ANNResult<U>,
S: FnOnce() -> ANNResult<U>,
{
match self.get_resumption_point(stage)? {
Some(_) => {
let result = operation()?;
self.update(Progress::Completed, next_stage)?;
Ok(result)
}
None => {
info!("[Stage:{:?}] Skip stage - invalid checkpoint", stage);
skip_handler()
}
}
}
}
pub trait CheckpointManagerClone {
fn clone_box(&self) -> Box<dyn CheckpointManager>;
}
impl<T> CheckpointManagerClone for T
where
T: 'static + CheckpointManager + Clone,
{
fn clone_box(&self) -> Box<dyn CheckpointManager> {
Box::new(self.clone())
}
}
#[cfg(test)]
mod tests {
use super::super::NaiveCheckpointRecordManager;
use super::*;
#[test]
fn test_checkpoint_manager_ext_execute_stage_with_resumption() {
let mut manager = NaiveCheckpointRecordManager;
let mut executed = false;
let result = manager.execute_stage(
WorkStage::Start,
WorkStage::End,
|| {
executed = true;
Ok(42)
},
|| Ok(0),
);
assert!(result.is_ok());
assert_eq!(result.unwrap(), 42);
assert!(executed);
}
#[test]
fn test_checkpoint_manager_clone_box() {
let manager = NaiveCheckpointRecordManager;
let boxed = manager.clone_box();
let result = boxed.get_resumption_point(WorkStage::Start);
assert!(result.is_ok());
assert_eq!(result.unwrap(), Some(0));
}
}