use std::path::Path;
use crate::backup::restore_path_from_latest_backup;
use crate::exit::FormatFailedError;
use crate::fallback::{EditError, EditErrorKind};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum PostWriteOnFailure {
#[default]
KeepWithError,
Revert,
}
#[derive(Debug, Clone, Default)]
pub struct PostWriteHooks {
pub format_cmd: Option<String>,
pub lint_cmd: Option<String>,
pub on_failure: PostWriteOnFailure,
pub timeout_secs: Option<u64>,
}
pub fn run_post_write_validation(
project_root: &Path,
path: &Path,
hooks: &PostWriteHooks,
) -> anyhow::Result<()> {
run_post_write_validation_with_session(project_root, path, hooks, None)
}
pub fn run_post_write_validation_with_session(
project_root: &Path,
path: &Path,
hooks: &PostWriteHooks,
backup_session: Option<&str>,
) -> anyhow::Result<()> {
let timeout = hooks.timeout_secs.unwrap_or(30);
let backup_root = path.parent().unwrap_or(project_root);
for (label, cmd) in [
("format", hooks.format_cmd.as_deref()),
("lint", hooks.lint_cmd.as_deref()),
] {
let Some(cmd) = cmd else {
continue;
};
if let Err(e) = run_hook_cmd(cmd, timeout, project_root, label) {
if hooks.on_failure == PostWriteOnFailure::Revert {
let restore = if let Some(ts) = backup_session {
crate::backup::restore_path_from_session(backup_root, ts, path)
} else {
restore_path_from_latest_backup(backup_root, path)
};
match restore {
Ok(true) => {}
Ok(false) => {
return Err(FormatFailedError::new(format!(
"{e}; also failed to revert {}: no backup session for path",
path.display()
))
.into());
}
Err(restore_err) => {
return Err(FormatFailedError::new(format!(
"{e}; also failed to revert {}: {restore_err}",
path.display()
))
.into());
}
}
}
return Err(e);
}
}
Ok(())
}
fn run_hook_cmd(cmd: &str, timeout_secs: u64, cwd: &Path, label: &str) -> anyhow::Result<()> {
let result = crate::exec::run_with_timeout(cmd, timeout_secs, cwd)
.map_err(|e| FormatFailedError::new(format!("{label} command failed ({cmd}): {e}")))?;
if !result.status.success() {
let stderr = if result.stderr_head.is_empty() {
String::new()
} else {
format!(": {}", result.stderr_head)
};
return Err(
FormatFailedError::new(format!("{label} command failed ({cmd}){stderr}")).into(),
);
}
Ok(())
}
pub trait PostWriteValidator {
fn validate(&self, path: &Path, before: &str, after: &str) -> Result<(), EditError>;
}
pub fn apply_post_write_validator<V: PostWriteValidator + ?Sized>(
project_root: &Path,
path: &Path,
before: &str,
after: &str,
validator: &V,
revert: bool,
) -> anyhow::Result<()> {
match validator.validate(path, before, after) {
Ok(()) => Ok(()),
Err(e) => {
if revert {
match restore_path_from_latest_backup(project_root, path) {
Ok(true) => {}
Ok(false) => {
return Err(EditError::new(
EditErrorKind::OperationFailed,
format!(
"{e}; also failed to revert {}: no backup session for path",
path.display()
),
)
.into());
}
Err(restore_err) => {
return Err(EditError::new(
EditErrorKind::OperationFailed,
format!(
"{e}; also failed to revert {}: {restore_err}",
path.display()
),
)
.into());
}
}
}
Err(e.into())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::{ApplyMode, ReplaceOptions, replace_text};
use crate::fallback::{EditErrorKind, edit_error_kind};
use std::fs;
use tempfile::TempDir;
#[test]
fn post_write_format_success_is_ok() {
let dir = TempDir::new().unwrap();
let file = dir.path().join("ok.txt");
fs::write(&file, "a\n").unwrap();
let hooks = PostWriteHooks {
format_cmd: Some("true".into()),
..Default::default()
};
run_post_write_validation(dir.path(), &file, &hooks).unwrap();
}
#[test]
fn post_write_format_failure_is_format_failed() {
let dir = TempDir::new().unwrap();
let file = dir.path().join("bad.txt");
fs::write(&file, "a\n").unwrap();
let hooks = PostWriteHooks {
format_cmd: Some("false".into()),
on_failure: PostWriteOnFailure::KeepWithError,
..Default::default()
};
let err = run_post_write_validation(dir.path(), &file, &hooks).unwrap_err();
assert!(
crate::exit::is_format_failed(&err),
"expected format_failed, got {err}"
);
assert_eq!(edit_error_kind(&err), Some(EditErrorKind::FormatFailed));
}
#[test]
fn post_write_revert_restores_path() {
let dir = TempDir::new().unwrap();
let file = dir.path().join("data.txt");
fs::write(&file, "before\n").unwrap();
let opts = ReplaceOptions {
require_change: true,
..Default::default()
};
replace_text(&file, "before", "after", &opts, ApplyMode::Apply, None).unwrap();
assert_eq!(fs::read_to_string(&file).unwrap(), "after\n");
let hooks = PostWriteHooks {
format_cmd: Some("false".into()),
on_failure: PostWriteOnFailure::Revert,
..Default::default()
};
let err = run_post_write_validation(dir.path(), &file, &hooks).unwrap_err();
assert!(crate::exit::is_format_failed(&err));
assert_eq!(
fs::read_to_string(&file).unwrap(),
"before\n",
"revert must restore pre-Apply bytes"
);
}
struct AlwaysFail;
impl PostWriteValidator for AlwaysFail {
fn validate(&self, _path: &Path, _before: &str, _after: &str) -> Result<(), EditError> {
Err(EditError::new(
EditErrorKind::OperationFailed,
"lint regression",
))
}
}
#[test]
fn post_write_validator_trait_reverts() {
let dir = TempDir::new().unwrap();
let file = dir.path().join("t.txt");
fs::write(&file, "old\n").unwrap();
replace_text(
&file,
"old",
"new",
&ReplaceOptions {
require_change: true,
..Default::default()
},
ApplyMode::Apply,
None,
)
.unwrap();
let err =
apply_post_write_validator(dir.path(), &file, "old\n", "new\n", &AlwaysFail, true)
.unwrap_err();
assert_eq!(edit_error_kind(&err), Some(EditErrorKind::OperationFailed));
assert_eq!(fs::read_to_string(&file).unwrap(), "old\n");
}
}