use core::fmt;
use std::collections::{BTreeMap, BTreeSet};
use std::io::Write as _;
use std::path::Path;
use serde::{Deserialize, Serialize};
use crate::file_lock::FileLock;
pub const OVERRIDES_VERSION: u32 = 1;
#[derive(Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct AppOverrides {
pub fields: serde_json::Map<String, serde_json::Value>,
pub declared: BTreeSet<String>,
pub declared_env: BTreeSet<String>,
}
impl fmt::Debug for AppOverrides {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AppOverrides")
.field("fields", &format_args!("<{} fields>", self.fields.len()))
.field("declared", &self.declared)
.field("declared_env", &self.declared_env)
.finish()
}
}
#[derive(Debug, Default, Serialize, Deserialize)]
struct OverridesFile {
version: u32,
apps: BTreeMap<String, AppOverrides>,
}
#[non_exhaustive]
#[derive(Debug)]
pub enum OverridesError {
Io(std::io::Error),
Decode(serde_json::Error),
FutureVersion(u32),
}
impl fmt::Display for OverridesError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io(err) => write!(f, "overrides store I/O failed: {err}"),
Self::Decode(err) => write!(f, "overrides store failed to parse: {err}"),
Self::FutureVersion(version) => {
write!(
f,
"overrides store is version {version}, newer than this build understands"
)
}
}
}
}
impl core::error::Error for OverridesError {
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match self {
Self::Io(err) => Some(err),
Self::Decode(err) => Some(err),
Self::FutureVersion(_) => None,
}
}
}
impl From<std::io::Error> for OverridesError {
fn from(source: std::io::Error) -> Self {
Self::Io(source)
}
}
impl From<serde_json::Error> for OverridesError {
fn from(source: serde_json::Error) -> Self {
Self::Decode(source)
}
}
fn read_file(path: &Path) -> Result<OverridesFile, OverridesError> {
let raw = match std::fs::read_to_string(path) {
Ok(raw) => raw,
Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
return Ok(OverridesFile::default());
}
Err(err) => return Err(OverridesError::Io(err)),
};
let file: OverridesFile = serde_json::from_str(&raw)?;
if file.version > OVERRIDES_VERSION {
return Err(OverridesError::FutureVersion(file.version));
}
Ok(file)
}
fn write_file(path: &Path, file: &OverridesFile) -> Result<(), OverridesError> {
let parent = path.parent().unwrap_or_else(|| Path::new("."));
let mut tmp = crate::atomic_file::create_staging_file(parent, "overrides", ".tmp")?;
let json = serde_json::to_string_pretty(file)?;
tmp.write_all(json.as_bytes())?;
tmp.write_all(b"\n")?;
tmp.as_file().sync_all()?;
tmp.persist(path)
.map_err(|err| OverridesError::Io(err.error))?;
crate::atomic_file::sync_dir(parent)?;
Ok(())
}
pub fn all(path: &Path) -> Result<BTreeMap<String, AppOverrides>, OverridesError> {
let _lock = FileLock::acquire(path)?;
Ok(read_file(path)?.apps)
}
pub fn get(path: &Path, name: &str) -> Result<Option<AppOverrides>, OverridesError> {
Ok(all(path)?.remove(name))
}
pub fn put(path: &Path, name: &str, value: &AppOverrides) -> Result<(), OverridesError> {
let _lock = FileLock::acquire(path)?;
let mut file = read_file(path)?;
file.version = OVERRIDES_VERSION;
file.apps.insert(name.to_string(), value.clone());
write_file(path, &file)
}
pub fn remove(path: &Path, name: &str) -> Result<bool, OverridesError> {
let _lock = FileLock::acquire(path)?;
let mut file = read_file(path)?;
let was_present = file.apps.remove(name).is_some();
if was_present {
file.version = OVERRIDES_VERSION;
write_file(path, &file)?;
}
Ok(was_present)
}
pub fn update(
path: &Path,
changes: &BTreeMap<String, Option<AppOverrides>>,
) -> Result<(), OverridesError> {
if changes.is_empty() {
return Ok(());
}
let _lock = FileLock::acquire(path)?;
let mut file = read_file(path)?;
for (name, change) in changes {
match change {
Some(value) => {
file.apps.insert(name.clone(), value.clone());
}
None => {
file.apps.remove(name);
}
}
}
file.version = OVERRIDES_VERSION;
write_file(path, &file)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn update_stores_removes_and_leaves_the_rest_alone() {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("overrides.json");
let record = |value: u64| AppOverrides {
fields: [("max_restarts".to_string(), serde_json::json!(value))]
.into_iter()
.collect(),
..AppOverrides::default()
};
put(&path, "web", &record(1)).unwrap();
put(&path, "worker", &record(2)).unwrap();
put(&path, "bystander", &record(3)).unwrap();
let changes = BTreeMap::from([
("web".to_string(), Some(record(9))),
("worker".to_string(), None),
]);
update(&path, &changes).unwrap();
let all = all(&path).unwrap();
assert_eq!(all.get("web"), Some(&record(9)));
assert_eq!(all.get("worker"), None);
assert_eq!(all.get("bystander"), Some(&record(3)));
}
#[test]
fn an_empty_update_writes_nothing() {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("overrides.json");
update(&path, &BTreeMap::new()).unwrap();
assert!(!path.exists(), "an empty batch created a store");
}
#[test]
fn put_then_get_round_trips() {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("overrides.json");
let mut fields = serde_json::Map::new();
fields.insert("max_memory".to_string(), serde_json::json!("512M"));
let value = AppOverrides {
fields,
declared: ["name", "script"].iter().map(|s| s.to_string()).collect(),
declared_env: BTreeSet::new(),
};
put(&path, "web", &value).unwrap();
assert_eq!(get(&path, "web").unwrap().as_ref(), Some(&value));
}
#[test]
fn a_missing_store_reads_as_empty() {
let dir = tempfile::TempDir::new().unwrap();
assert!(all(&dir.path().join("overrides.json")).unwrap().is_empty());
}
#[cfg(unix)]
#[test]
fn the_store_is_owner_only() {
use std::os::unix::fs::PermissionsExt as _;
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("overrides.json");
put(&path, "web", &AppOverrides::default()).unwrap();
let mode = std::fs::metadata(&path).unwrap().permissions().mode();
assert_eq!(mode & 0o777, 0o600, "mode was {:o}", mode & 0o777);
}
#[test]
fn debug_redacts_override_values() {
let mut fields = serde_json::Map::new();
fields.insert(
"env".to_string(),
serde_json::json!({"DATABASE_URL": "postgres://hunter2"}),
);
let value = AppOverrides {
fields,
..AppOverrides::default()
};
let rendered = format!("{value:?}");
assert!(!rendered.contains("hunter2"), "leaked: {rendered}");
assert_eq!(
rendered,
"AppOverrides { fields: <1 fields>, declared: {}, declared_env: {} }"
);
}
#[test]
fn a_future_version_refuses_without_clobbering() {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("overrides.json");
std::fs::write(&path, r#"{"version":99,"apps":{}}"#).unwrap();
assert!(matches!(
get(&path, "web"),
Err(OverridesError::FutureVersion(99))
));
assert_eq!(
std::fs::read_to_string(&path).unwrap(),
r#"{"version":99,"apps":{}}"#
);
}
#[test]
fn two_concurrent_writers_lose_nothing() {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("overrides.json");
const PER_WRITER: usize = 50;
let (done_tx, done_rx) = std::sync::mpsc::channel();
for writer in 0..2 {
let path = path.clone();
let done_tx = done_tx.clone();
std::thread::spawn(move || {
for n in 0..PER_WRITER {
put(&path, &format!("w{writer}-{n}"), &AppOverrides::default()).unwrap();
}
done_tx.send(()).unwrap();
});
}
drop(done_tx);
for _ in 0..2 {
done_rx
.recv_timeout(std::time::Duration::from_secs(60))
.expect("a writer did not finish within 60s");
}
assert_eq!(all(&path).unwrap().len(), PER_WRITER * 2);
}
}