use std::fs::File;
use std::io::Write as IoWrite;
use std::path::{Path, PathBuf};
use std::time::{SystemTime, UNIX_EPOCH};
struct TempFileGuard {
path: PathBuf,
armed: bool,
}
impl TempFileGuard {
fn new(path: PathBuf) -> Self {
Self { path, armed: true }
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for TempFileGuard {
fn drop(&mut self) {
if self.armed {
let _ = std::fs::remove_file(&self.path);
}
}
}
fn temp_path_for(path: &Path) -> PathBuf {
let parent = path.parent().unwrap_or_else(|| Path::new("."));
let file = path
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("shard.rkyv");
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
parent.join(format!("{}.tmp.{}.{}", file, std::process::id(), nanos))
}
pub fn write_bytes_atomic(path: &Path, bytes: &[u8]) -> Result<(), String> {
let parent = path.parent().unwrap_or_else(|| Path::new("."));
let _ = std::fs::create_dir_all(parent);
let tmp_path = temp_path_for(path);
let mut guard = TempFileGuard::new(tmp_path.clone());
{
let mut f = File::create(&tmp_path).map_err(|e| e.to_string())?;
f.write_all(bytes).map_err(|e| e.to_string())?;
f.sync_all().map_err(|e| e.to_string())?;
}
std::fs::rename(&tmp_path, path).map_err(|e| e.to_string())?;
guard.disarm();
let reaped = reap_orphan_temps(path);
if reaped > 0 {
tracing::info!(
path = %path.display(),
reaped,
"shard write: removed temp files left by dead processes"
);
}
Ok(())
}
pub fn reap_orphan_temps(path: &Path) -> usize {
let Some(parent) = path.parent() else {
return 0;
};
let Some(file) = path.file_name().and_then(|s| s.to_str()) else {
return 0;
};
let prefix = format!("{}.tmp.", file);
let me = std::process::id() as i32;
let Ok(dir) = std::fs::read_dir(parent) else {
return 0;
};
let mut removed = 0usize;
for entry in dir.flatten() {
let Ok(name) = entry.file_name().into_string() else {
continue;
};
let Some(rest) = name.strip_prefix(&prefix) else {
continue;
};
let Some((pid_str, nanos)) = rest.split_once('.') else {
continue;
};
if nanos.is_empty() || !nanos.bytes().all(|b| b.is_ascii_digit()) {
continue;
}
let Ok(pid) = pid_str.parse::<i32>() else {
continue;
};
if pid <= 0 || pid == me || pid_is_alive(pid) {
continue;
}
if std::fs::remove_file(entry.path()).is_ok() {
removed += 1;
}
}
removed
}
fn pid_is_alive(pid: i32) -> bool {
match nix::sys::signal::kill(nix::unistd::Pid::from_raw(pid), None) {
Ok(()) => true,
Err(nix::errno::Errno::EPERM) => true,
Err(_) => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn a_successful_write_leaves_no_temp_behind() {
let dir = tempdir().unwrap();
let path = dir.path().join("shard.rkyv");
write_bytes_atomic(&path, b"payload").unwrap();
assert_eq!(std::fs::read(&path).unwrap(), b"payload");
let leftovers: Vec<_> = std::fs::read_dir(dir.path())
.unwrap()
.flatten()
.map(|e| e.file_name().into_string().unwrap())
.filter(|n| n.contains(".tmp."))
.collect();
assert!(leftovers.is_empty(), "temp files left behind: {leftovers:?}");
}
#[test]
fn a_dead_pids_temp_is_reaped_and_a_live_pids_is_not() {
let dir = tempdir().unwrap();
let path = dir.path().join("shard.rkyv");
let live = dir.path().join("shard.rkyv.tmp.1.123");
let dead = dir.path().join("shard.rkyv.tmp.2147483646.456");
let other = dir.path().join("other.rkyv.tmp.2147483646.789");
std::fs::write(&live, b"x").unwrap();
std::fs::write(&dead, b"x").unwrap();
std::fs::write(&other, b"x").unwrap();
assert_eq!(reap_orphan_temps(&path), 1);
assert!(live.exists(), "a live process's temp was deleted");
assert!(!dead.exists(), "a dead process's temp was not reaped");
assert!(other.exists(), "an unrelated file was deleted");
}
}