use std::path::PathBuf;
use crate::executor::CheckpointEntryView;
use super::transaction::TransactionSnapshot;
#[derive(Debug)]
pub(crate) struct Checkpoint {
pub(crate) before_snapshot: TransactionSnapshot,
pub(crate) command: String,
pub(crate) paths: Vec<PathBuf>,
pub(crate) captured_at_secs: u64,
}
#[derive(Debug)]
pub(crate) struct CheckpointStack {
pub(crate) undo: Vec<Checkpoint>,
pub(crate) redo: Vec<Checkpoint>,
pub(crate) max_checkpoints: usize,
}
impl CheckpointStack {
pub(crate) fn new(max_checkpoints: usize) -> Self {
Self {
undo: Vec::new(),
redo: Vec::new(),
max_checkpoints,
}
}
pub(crate) fn record(&mut self, checkpoint: Checkpoint) {
self.redo.clear();
if self.max_checkpoints > 0 && self.undo.len() >= self.max_checkpoints {
self.undo.remove(0);
}
self.undo.push(checkpoint);
}
pub(crate) fn undo(&mut self, n: usize, max_snapshot_bytes: u64) -> UndoRedoResult {
if n == 0 {
return UndoRedoResult {
reverted_commands: 0,
restored: 0,
deleted: 0,
message: "Undo count must be > 0.".to_owned(),
};
}
let available = self.undo.len();
let count = n.min(available);
if count == 0 {
return UndoRedoResult {
reverted_commands: 0,
restored: 0,
deleted: 0,
message: "Nothing to undo.".to_owned(),
};
}
let mut total_restored = 0usize;
let mut total_deleted = 0usize;
let mut reverted = Vec::new();
for _ in 0..count {
let Some(cp) = self.undo.pop() else {
break;
};
let redo_snap = match TransactionSnapshot::capture(&cp.paths, max_snapshot_bytes) {
Ok(s) => Some(s),
Err(e) => {
tracing::warn!(err = %e, "checkpoint undo: redo re-capture failed, redo entry skipped");
None
}
};
let cmd = cp.command.clone();
let paths = cp.paths.clone();
let captured_at_secs = cp.captured_at_secs;
match cp.before_snapshot.rollback() {
Ok(report) => {
total_restored += report.restored_count;
total_deleted += report.deleted_count;
reverted.push(cmd.clone());
if let Some(snap) = redo_snap {
self.redo.push(Checkpoint {
before_snapshot: snap,
command: cmd,
paths,
captured_at_secs,
});
}
}
Err(e) => {
tracing::error!(err = %e, cmd = %cmd, "checkpoint undo: rollback failed");
}
}
}
let actual = reverted.len();
let message = if actual == 0 {
"Undo failed: rollback errors for all checkpoints.".to_owned()
} else if available > count {
format!(
"Undid {} of {} available command(s); {} more available.",
actual,
available,
available - count
)
} else {
format!("Undid {actual} command(s).")
};
UndoRedoResult {
reverted_commands: actual,
restored: total_restored,
deleted: total_deleted,
message,
}
}
pub(crate) fn redo(&mut self, max_snapshot_bytes: u64) -> UndoRedoResult {
let Some(cp) = self.redo.pop() else {
return UndoRedoResult {
reverted_commands: 0,
restored: 0,
deleted: 0,
message: "Nothing to redo.".to_owned(),
};
};
let undo_snap = match TransactionSnapshot::capture(&cp.paths, max_snapshot_bytes) {
Ok(s) => Some(s),
Err(e) => {
tracing::warn!(err = %e, "checkpoint redo: undo re-capture failed, undo entry skipped");
None
}
};
let cmd = cp.command.clone();
let paths = cp.paths.clone();
let captured_at_secs = cp.captured_at_secs;
match cp.before_snapshot.rollback() {
Ok(report) => {
if let Some(snap) = undo_snap {
self.undo.push(Checkpoint {
before_snapshot: snap,
command: cmd.clone(),
paths,
captured_at_secs,
});
}
UndoRedoResult {
reverted_commands: 1,
restored: report.restored_count,
deleted: report.deleted_count,
message: format!("Redone: {cmd}"),
}
}
Err(e) => {
tracing::error!(err = %e, cmd = %cmd, "checkpoint redo: rollback failed");
UndoRedoResult {
reverted_commands: 0,
restored: 0,
deleted: 0,
message: format!("Redo failed: {e}"),
}
}
}
}
pub(crate) fn list_undo(&self) -> Vec<CheckpointEntryView> {
self.undo
.iter()
.enumerate()
.rev()
.map(|(stack_idx, cp)| CheckpointEntryView {
index: self.undo.len() - 1 - stack_idx,
command: cp.command.clone(),
captured_at_secs: cp.captured_at_secs,
file_count: cp.paths.len(),
})
.collect()
}
pub(crate) fn redo_depth(&self) -> usize {
self.redo.len()
}
}
#[derive(Debug)]
pub(crate) struct UndoRedoResult {
pub(crate) reverted_commands: usize,
pub(crate) restored: usize,
pub(crate) deleted: usize,
pub(crate) message: String,
}
#[cfg(test)]
mod tests {
use std::path::PathBuf;
use super::*;
fn make_checkpoint(paths: &[PathBuf]) -> Checkpoint {
let snap = TransactionSnapshot::capture(paths, 0).unwrap();
Checkpoint {
before_snapshot: snap,
command: "echo test".to_owned(),
paths: paths.to_vec(),
captured_at_secs: 0,
}
}
#[test]
fn undo_zero_count_on_non_empty_stack() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("f.txt");
std::fs::write(&p, "v0").unwrap();
let mut stack = CheckpointStack::new(10);
stack.record(make_checkpoint(std::slice::from_ref(&p)));
assert_eq!(stack.undo.len(), 1);
let r = stack.undo(0, 0);
assert_eq!(r.message, "Undo count must be > 0.");
assert_eq!(r.reverted_commands, 0);
assert_eq!(stack.undo.len(), 1);
}
#[test]
fn undo_empty_stack_returns_nothing_to_undo() {
let mut stack = CheckpointStack::new(10);
let r = stack.undo(1, 0);
assert_eq!(r.reverted_commands, 0);
assert_eq!(r.message, "Nothing to undo.");
}
#[test]
fn redo_empty_stack_returns_nothing_to_redo() {
let mut stack = CheckpointStack::new(10);
let r = stack.redo(0);
assert_eq!(r.reverted_commands, 0);
assert_eq!(r.message, "Nothing to redo.");
}
#[test]
fn record_at_cap_evicts_oldest() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("f.txt");
std::fs::write(&p, "v0").unwrap();
let mut stack = CheckpointStack::new(2);
let cp0 = make_checkpoint(std::slice::from_ref(&p));
let cp1 = make_checkpoint(std::slice::from_ref(&p));
let cp2 = make_checkpoint(std::slice::from_ref(&p));
stack.record(cp0);
stack.record(cp1);
stack.record(cp2);
assert_eq!(stack.undo.len(), 2);
}
#[test]
fn record_clears_redo_stack() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("f.txt");
std::fs::write(&p, "v0").unwrap();
let mut stack = CheckpointStack::new(10);
stack.record(make_checkpoint(std::slice::from_ref(&p)));
stack.undo(1, 0); assert_eq!(stack.redo.len(), 1);
std::fs::write(&p, "v1").unwrap();
stack.record(make_checkpoint(std::slice::from_ref(&p)));
assert_eq!(stack.redo.len(), 0);
}
#[test]
fn undo_n_greater_than_available_clamped() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("f.txt");
std::fs::write(&p, "v0").unwrap();
let mut stack = CheckpointStack::new(10);
stack.record(make_checkpoint(std::slice::from_ref(&p)));
let r = stack.undo(99, 0);
assert_eq!(r.reverted_commands, 1);
}
#[test]
fn list_undo_most_recent_first() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("f.txt");
std::fs::write(&p, "v0").unwrap();
let mut stack = CheckpointStack::new(10);
let mut cp_a = make_checkpoint(std::slice::from_ref(&p));
cp_a.command = "cmd_a".to_owned();
cp_a.captured_at_secs = 1;
let mut cp_b = make_checkpoint(std::slice::from_ref(&p));
cp_b.command = "cmd_b".to_owned();
cp_b.captured_at_secs = 2;
stack.record(cp_a);
stack.record(cp_b);
let list = stack.list_undo();
assert_eq!(list[0].command, "cmd_b");
assert_eq!(list[1].command, "cmd_a");
}
#[test]
fn undo_redo_roundtrip_restores_content() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("target.txt");
std::fs::write(&p, "before").unwrap();
let mut stack = CheckpointStack::new(10);
stack.record(make_checkpoint(std::slice::from_ref(&p)));
std::fs::write(&p, "after").unwrap();
let r = stack.undo(1, 0);
assert_eq!(r.reverted_commands, 1);
assert_eq!(std::fs::read_to_string(&p).unwrap(), "before");
let r2 = stack.redo(0);
assert_eq!(r2.reverted_commands, 1);
assert_eq!(std::fs::read_to_string(&p).unwrap(), "after");
}
#[test]
fn redo_of_delete_removes_file_again() {
let dir = tempfile::tempdir().unwrap();
let p = dir.path().join("to_delete.txt");
std::fs::write(&p, "exists").unwrap();
let mut stack = CheckpointStack::new(10);
stack.record(make_checkpoint(std::slice::from_ref(&p)));
std::fs::remove_file(&p).unwrap();
stack.undo(1, 0);
assert!(p.exists(), "undo must restore the deleted file");
stack.redo(0);
assert!(!p.exists(), "redo must re-delete the restored file");
}
}