mermaid_runtime/
atomic.rs1use std::fs::{self, File};
12use std::io::Write;
13use std::path::Path;
14use std::sync::atomic::{AtomicU64, Ordering};
15use std::time::{Duration, SystemTime};
16
17static COUNTER: AtomicU64 = AtomicU64::new(0);
18
19const STALE_TEMP_SECS: u64 = 3600;
25
26pub fn write_atomic(path: &Path, bytes: &[u8]) -> std::io::Result<()> {
37 write_atomic_inner(path, bytes, None)
38}
39
40pub fn write_atomic_with_mode(path: &Path, bytes: &[u8], mode: u32) -> std::io::Result<()> {
50 write_atomic_inner(path, bytes, Some(mode))
51}
52
53fn write_atomic_inner(path: &Path, bytes: &[u8], mode: Option<u32>) -> std::io::Result<()> {
54 let parent = path.parent().unwrap_or_else(|| Path::new("."));
55 fs::create_dir_all(parent)?;
56
57 let stem = path.file_name().and_then(|n| n.to_str()).unwrap_or("tmp");
58
59 sweep_stale_temps(parent, stem, Duration::from_secs(STALE_TEMP_SECS));
66
67 let n = COUNTER.fetch_add(1, Ordering::Relaxed);
68 let tmp = parent.join(format!(".{}.{}.{}.tmp", stem, std::process::id(), n));
69
70 {
71 let mut f = create_temp(&tmp, mode)?;
72 f.write_all(bytes)?;
73 f.sync_all()?;
74 }
75
76 if let Err(e) = fs::rename(&tmp, path) {
77 let _ = fs::remove_file(&tmp);
78 return Err(e);
79 }
80
81 if let Ok(dir) = File::open(parent) {
84 let _ = dir.sync_all();
85 }
86 Ok(())
87}
88
89#[cfg(unix)]
92fn create_temp(tmp: &Path, mode: Option<u32>) -> std::io::Result<File> {
93 match mode {
94 Some(mode) => {
95 use std::os::unix::fs::OpenOptionsExt;
96 std::fs::OpenOptions::new()
97 .write(true)
98 .create(true)
99 .truncate(true)
100 .mode(mode)
101 .open(tmp)
102 },
103 None => File::create(tmp),
104 }
105}
106
107#[cfg(not(unix))]
108fn create_temp(tmp: &Path, _mode: Option<u32>) -> std::io::Result<File> {
109 File::create(tmp)
110}
111
112fn sweep_stale_temps(parent: &Path, stem: &str, max_age: Duration) {
123 let prefix = format!(".{stem}.");
124 let Ok(entries) = fs::read_dir(parent) else {
125 return;
126 };
127 let now = SystemTime::now();
128 for entry in entries.flatten() {
129 let name = entry.file_name();
130 let Some(name) = name.to_str() else {
131 continue;
132 };
133 if !name.starts_with(&prefix) || !name.ends_with(".tmp") {
134 continue;
135 }
136 let stale = entry
137 .metadata()
138 .and_then(|m| m.modified())
139 .ok()
140 .and_then(|mtime| now.duration_since(mtime).ok())
141 .map(|age| age >= max_age)
142 .unwrap_or(false);
143 if stale {
144 let _ = fs::remove_file(entry.path());
145 }
146 }
147}
148
149#[cfg(test)]
150mod tests {
151 use super::*;
152
153 #[test]
154 fn atomic_write_replaces_existing_and_no_temp_left() {
155 let dir = std::env::temp_dir().join(format!("mermaid_atomic_{}", std::process::id()));
156 let _ = fs::create_dir_all(&dir);
157 let target = dir.join("conv.json");
158 write_atomic(&target, b"first").unwrap();
159 assert_eq!(fs::read_to_string(&target).unwrap(), "first");
160 write_atomic(&target, b"second").unwrap();
161 assert_eq!(fs::read_to_string(&target).unwrap(), "second");
162 let leftovers = fs::read_dir(&dir)
164 .unwrap()
165 .flatten()
166 .filter(|e| e.file_name().to_string_lossy().ends_with(".tmp"))
167 .count();
168 assert_eq!(leftovers, 0);
169 let _ = fs::remove_dir_all(&dir);
170 }
171
172 #[test]
173 fn sweep_removes_only_matching_stale_temps() {
174 let dir = std::env::temp_dir().join(format!("mermaid_atomic_sweep_{}", std::process::id()));
175 let _ = fs::remove_dir_all(&dir);
176 let _ = fs::create_dir_all(&dir);
177
178 let target = dir.join("conv.json");
179 write_atomic(&target, b"live").unwrap();
180
181 let orphan = dir.join(".conv.json.99999.0.tmp");
183 fs::write(&orphan, b"half-written").unwrap();
184 let other = dir.join(".other.json.99999.0.tmp");
186 fs::write(&other, b"someone else").unwrap();
187 let unrelated = dir.join("notes.txt");
189 fs::write(&unrelated, b"keep me").unwrap();
190
191 sweep_stale_temps(&dir, "conv.json", Duration::ZERO);
193
194 assert!(!orphan.exists(), "matching stale temp must be swept");
195 assert!(other.exists(), "a different target's temp must survive");
196 assert!(unrelated.exists(), "unrelated files must survive");
197 assert!(target.exists(), "the destination must never be swept");
198 assert_eq!(fs::read_to_string(&target).unwrap(), "live");
199
200 let _ = fs::remove_dir_all(&dir);
201 }
202
203 #[test]
204 fn sweep_preserves_fresh_in_flight_temps() {
205 let dir = std::env::temp_dir().join(format!("mermaid_atomic_fresh_{}", std::process::id()));
206 let _ = fs::remove_dir_all(&dir);
207 let _ = fs::create_dir_all(&dir);
208
209 let fresh = dir.join(".conv.json.12345.7.tmp");
211 fs::write(&fresh, b"being written").unwrap();
212
213 sweep_stale_temps(&dir, "conv.json", Duration::from_secs(STALE_TEMP_SECS));
215
216 assert!(fresh.exists(), "a fresh/in-flight temp must not be swept");
217 let _ = fs::remove_dir_all(&dir);
218 }
219
220 #[test]
221 #[cfg(unix)]
222 fn write_atomic_with_mode_creates_0600_and_leaves_no_temp() {
223 use std::os::unix::fs::PermissionsExt;
224 let dir = std::env::temp_dir().join(format!("mermaid_atomic_mode_{}", std::process::id()));
225 let _ = fs::remove_dir_all(&dir);
226 let _ = fs::create_dir_all(&dir);
227 let target = dir.join("config.toml");
228
229 write_atomic_with_mode(&target, b"secret = true", 0o600).unwrap();
230 assert_eq!(fs::read_to_string(&target).unwrap(), "secret = true");
231 let mode = fs::metadata(&target).unwrap().permissions().mode() & 0o777;
232 assert_eq!(mode, 0o600, "config must be created 0600, not at umask");
233
234 fs::set_permissions(&target, fs::Permissions::from_mode(0o644)).unwrap();
237 write_atomic_with_mode(&target, b"secret = false", 0o600).unwrap();
238 let mode = fs::metadata(&target).unwrap().permissions().mode() & 0o777;
239 assert_eq!(
240 mode, 0o600,
241 "an overwrite must not leave the file world-readable"
242 );
243
244 let leftovers = fs::read_dir(&dir)
245 .unwrap()
246 .flatten()
247 .filter(|e| e.file_name().to_string_lossy().ends_with(".tmp"))
248 .count();
249 assert_eq!(leftovers, 0, "no temp left behind");
250
251 let _ = fs::remove_dir_all(&dir);
252 }
253}