1use core::fmt;
15use std::collections::{BTreeMap, BTreeSet};
16use std::io::Write as _;
17use std::path::Path;
18#[cfg(any(unix, windows))]
21use std::path::PathBuf;
22
23use serde::{Deserialize, Serialize};
24
25pub const OVERRIDES_VERSION: u32 = 1;
31
32#[derive(Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
41pub struct AppOverrides {
42 pub fields: serde_json::Map<String, serde_json::Value>,
45 pub declared: BTreeSet<String>,
47 pub declared_env: BTreeSet<String>,
49}
50
51impl fmt::Debug for AppOverrides {
54 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
55 f.debug_struct("AppOverrides")
56 .field("fields", &format_args!("<{} fields>", self.fields.len()))
57 .field("declared", &self.declared)
58 .field("declared_env", &self.declared_env)
59 .finish()
60 }
61}
62
63#[derive(Debug, Default, Serialize, Deserialize)]
68struct OverridesFile {
69 version: u32,
70 apps: BTreeMap<String, AppOverrides>,
71}
72
73#[non_exhaustive]
83#[derive(Debug)]
84pub enum OverridesError {
85 Io(std::io::Error),
87 Decode(serde_json::Error),
93 FutureVersion(u32),
96}
97
98impl fmt::Display for OverridesError {
99 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
100 match self {
101 Self::Io(err) => write!(f, "overrides store I/O failed: {err}"),
102 Self::Decode(err) => write!(f, "overrides store failed to parse: {err}"),
103 Self::FutureVersion(version) => {
104 write!(
105 f,
106 "overrides store is version {version}, newer than this build understands"
107 )
108 }
109 }
110 }
111}
112
113impl core::error::Error for OverridesError {
114 fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
115 match self {
116 Self::Io(err) => Some(err),
117 Self::Decode(err) => Some(err),
118 Self::FutureVersion(_) => None,
119 }
120 }
121}
122
123impl From<std::io::Error> for OverridesError {
124 fn from(source: std::io::Error) -> Self {
125 Self::Io(source)
126 }
127}
128
129impl From<serde_json::Error> for OverridesError {
130 fn from(source: serde_json::Error) -> Self {
131 Self::Decode(source)
132 }
133}
134
135#[cfg(any(unix, windows))]
142fn lock_path(path: &Path) -> PathBuf {
143 let mut name = path
144 .file_name()
145 .map(std::ffi::OsStr::to_os_string)
146 .unwrap_or_default();
147 name.push(".lock");
148 path.parent().unwrap_or_else(|| Path::new(".")).join(name)
149}
150
151struct OverridesLock {
157 #[cfg(unix)]
160 _flock: nix::fcntl::Flock<std::fs::File>,
161 #[cfg(windows)]
167 _handle: std::fs::File,
168}
169
170impl OverridesLock {
171 #[cfg(unix)]
178 fn acquire(path: &Path) -> std::io::Result<Self> {
179 use nix::fcntl::{Flock, FlockArg};
180 use std::os::unix::fs::OpenOptionsExt as _;
181
182 let file = std::fs::OpenOptions::new()
183 .write(true)
184 .create(true)
185 .truncate(false)
186 .mode(crate::atomic_file::OWNER_ONLY_FILE_MODE)
187 .open(lock_path(path))?;
188
189 Flock::lock(file, FlockArg::LockExclusive)
190 .map(|flock| Self { _flock: flock })
191 .map_err(|(_file, errno)| std::io::Error::from(errno))
192 }
193
194 #[cfg(windows)]
206 fn acquire(path: &Path) -> std::io::Result<Self> {
207 use std::os::windows::fs::OpenOptionsExt as _;
208
209 const ERROR_SHARING_VIOLATION: i32 = 32;
213
214 const RETRY_INTERVAL: std::time::Duration = std::time::Duration::from_millis(2);
219
220 let lock_path = lock_path(path);
221 loop {
222 match std::fs::OpenOptions::new()
223 .write(true)
224 .create(true)
225 .truncate(false)
226 .share_mode(0)
227 .open(&lock_path)
228 {
229 Ok(handle) => return Ok(Self { _handle: handle }),
230 Err(error) if error.raw_os_error() == Some(ERROR_SHARING_VIOLATION) => {
231 std::thread::sleep(RETRY_INTERVAL);
232 }
233 Err(error) => return Err(error),
234 }
235 }
236 }
237}
238
239fn read_file(path: &Path) -> Result<OverridesFile, OverridesError> {
245 let raw = match std::fs::read_to_string(path) {
246 Ok(raw) => raw,
247 Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
248 return Ok(OverridesFile::default());
249 }
250 Err(err) => return Err(OverridesError::Io(err)),
251 };
252 let file: OverridesFile = serde_json::from_str(&raw)?;
253 if file.version > OVERRIDES_VERSION {
254 return Err(OverridesError::FutureVersion(file.version));
255 }
256 Ok(file)
257}
258
259fn write_file(path: &Path, file: &OverridesFile) -> Result<(), OverridesError> {
262 let parent = path.parent().unwrap_or_else(|| Path::new("."));
263 let mut tmp = crate::atomic_file::create_staging_file(parent, "overrides", ".tmp")?;
264
265 let json = serde_json::to_string_pretty(file)?;
266 tmp.write_all(json.as_bytes())?;
267 tmp.write_all(b"\n")?;
268 tmp.as_file().sync_all()?;
269
270 tmp.persist(path)
274 .map_err(|err| OverridesError::Io(err.error))?;
275
276 crate::atomic_file::sync_dir(parent)?;
279 Ok(())
280}
281
282pub fn all(path: &Path) -> Result<BTreeMap<String, AppOverrides>, OverridesError> {
293 let _lock = OverridesLock::acquire(path)?;
296 Ok(read_file(path)?.apps)
297}
298
299pub fn get(path: &Path, name: &str) -> Result<Option<AppOverrides>, OverridesError> {
306 Ok(all(path)?.remove(name))
307}
308
309pub fn put(path: &Path, name: &str, value: &AppOverrides) -> Result<(), OverridesError> {
319 let _lock = OverridesLock::acquire(path)?;
320 let mut file = read_file(path)?;
321 file.version = OVERRIDES_VERSION;
322 file.apps.insert(name.to_string(), value.clone());
323 write_file(path, &file)
324}
325
326pub fn remove(path: &Path, name: &str) -> Result<bool, OverridesError> {
332 let _lock = OverridesLock::acquire(path)?;
333 let mut file = read_file(path)?;
334 let was_present = file.apps.remove(name).is_some();
335 if was_present {
336 file.version = OVERRIDES_VERSION;
337 write_file(path, &file)?;
338 }
339 Ok(was_present)
340}
341
342pub fn update(
355 path: &Path,
356 changes: &BTreeMap<String, Option<AppOverrides>>,
357) -> Result<(), OverridesError> {
358 if changes.is_empty() {
359 return Ok(());
360 }
361 let _lock = OverridesLock::acquire(path)?;
362 let mut file = read_file(path)?;
363 for (name, change) in changes {
364 match change {
365 Some(value) => {
366 file.apps.insert(name.clone(), value.clone());
367 }
368 None => {
369 file.apps.remove(name);
370 }
371 }
372 }
373 file.version = OVERRIDES_VERSION;
374 write_file(path, &file)
375}
376
377#[cfg(test)]
378mod tests {
379 use super::*;
380
381 #[test]
382 fn update_stores_removes_and_leaves_the_rest_alone() {
383 let dir = tempfile::TempDir::new().unwrap();
384 let path = dir.path().join("overrides.json");
385 let record = |value: u64| AppOverrides {
386 fields: [("max_restarts".to_string(), serde_json::json!(value))]
387 .into_iter()
388 .collect(),
389 ..AppOverrides::default()
390 };
391 put(&path, "web", &record(1)).unwrap();
392 put(&path, "worker", &record(2)).unwrap();
393 put(&path, "bystander", &record(3)).unwrap();
394
395 let changes = BTreeMap::from([
396 ("web".to_string(), Some(record(9))),
397 ("worker".to_string(), None),
398 ]);
399 update(&path, &changes).unwrap();
400
401 let all = all(&path).unwrap();
402 assert_eq!(all.get("web"), Some(&record(9)));
403 assert_eq!(all.get("worker"), None);
404 assert_eq!(all.get("bystander"), Some(&record(3)));
405 }
406
407 #[test]
408 fn an_empty_update_writes_nothing() {
409 let dir = tempfile::TempDir::new().unwrap();
410 let path = dir.path().join("overrides.json");
411 update(&path, &BTreeMap::new()).unwrap();
412 assert!(!path.exists(), "an empty batch created a store");
413 }
414
415 #[test]
416 fn put_then_get_round_trips() {
417 let dir = tempfile::TempDir::new().unwrap();
418 let path = dir.path().join("overrides.json");
419 let mut fields = serde_json::Map::new();
420 fields.insert("max_memory".to_string(), serde_json::json!("512M"));
421 let value = AppOverrides {
422 fields,
423 declared: ["name", "script"].iter().map(|s| s.to_string()).collect(),
424 declared_env: BTreeSet::new(),
425 };
426 put(&path, "web", &value).unwrap();
427 assert_eq!(get(&path, "web").unwrap().as_ref(), Some(&value));
428 }
429
430 #[test]
431 fn a_missing_store_reads_as_empty() {
432 let dir = tempfile::TempDir::new().unwrap();
433 assert!(all(&dir.path().join("overrides.json")).unwrap().is_empty());
434 }
435
436 #[cfg(unix)]
438 #[test]
439 fn the_store_is_owner_only() {
440 use std::os::unix::fs::PermissionsExt as _;
441 let dir = tempfile::TempDir::new().unwrap();
442 let path = dir.path().join("overrides.json");
443 put(&path, "web", &AppOverrides::default()).unwrap();
444 let mode = std::fs::metadata(&path).unwrap().permissions().mode();
445 assert_eq!(mode & 0o777, 0o600, "mode was {:o}", mode & 0o777);
446 }
447
448 #[test]
449 fn debug_redacts_override_values() {
450 let mut fields = serde_json::Map::new();
451 fields.insert(
452 "env".to_string(),
453 serde_json::json!({"DATABASE_URL": "postgres://hunter2"}),
454 );
455 let value = AppOverrides {
456 fields,
457 ..AppOverrides::default()
458 };
459 let rendered = format!("{value:?}");
460 assert!(!rendered.contains("hunter2"), "leaked: {rendered}");
461 assert_eq!(
464 rendered,
465 "AppOverrides { fields: <1 fields>, declared: {}, declared_env: {} }"
466 );
467 }
468
469 #[test]
470 fn a_future_version_refuses_without_clobbering() {
471 let dir = tempfile::TempDir::new().unwrap();
472 let path = dir.path().join("overrides.json");
473 std::fs::write(&path, r#"{"version":99,"apps":{}}"#).unwrap();
474 assert!(matches!(
475 get(&path, "web"),
476 Err(OverridesError::FutureVersion(99))
477 ));
478 assert_eq!(
479 std::fs::read_to_string(&path).unwrap(),
480 r#"{"version":99,"apps":{}}"#
481 );
482 }
483
484 #[test]
487 fn two_concurrent_writers_lose_nothing() {
488 let dir = tempfile::TempDir::new().unwrap();
489 let path = dir.path().join("overrides.json");
490 const PER_WRITER: usize = 50;
491
492 let (done_tx, done_rx) = std::sync::mpsc::channel();
493 for writer in 0..2 {
494 let path = path.clone();
495 let done_tx = done_tx.clone();
496 std::thread::spawn(move || {
497 for n in 0..PER_WRITER {
498 put(&path, &format!("w{writer}-{n}"), &AppOverrides::default()).unwrap();
499 }
500 done_tx.send(()).unwrap();
501 });
502 }
503 drop(done_tx);
504 for _ in 0..2 {
505 done_rx
506 .recv_timeout(std::time::Duration::from_secs(60))
507 .expect("a writer did not finish within 60s");
508 }
509
510 assert_eq!(all(&path).unwrap().len(), PER_WRITER * 2);
511 }
512}