use std::path::{Path, PathBuf};
use serde::Serialize;
use serde::de::DeserializeOwned;
pub const TRUSTY_TOOLS_DIR: &str = ".trusty-tools";
pub const CONFIG_FILE: &str = "config.yaml";
#[derive(Debug, thiserror::Error)]
pub enum ConfigError {
#[error("config I/O error at {path}: {source}")]
Io {
path: PathBuf,
source: std::io::Error,
},
#[error("config YAML error at {path}: {detail}")]
Yaml {
path: PathBuf,
detail: YamlErrorDetail,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum YamlErrorKind {
InvalidType,
InvalidValue,
InvalidLength,
UnknownVariant,
UnknownField,
MissingField,
DuplicateField,
Malformed,
Serialize,
}
impl YamlErrorKind {
const MARKERS: [(&'static str, Self); 7] = [
("invalid type: ", Self::InvalidType),
("invalid value: ", Self::InvalidValue),
("invalid length ", Self::InvalidLength),
("unknown variant `", Self::UnknownVariant),
("unknown field `", Self::UnknownField),
("missing field `", Self::MissingField),
("duplicate field `", Self::DuplicateField),
];
fn classify(message: &str) -> Self {
if message.contains("unknown field `") {
return Self::UnknownField;
}
Self::MARKERS
.iter()
.filter_map(|(marker, kind)| message.find(marker).map(|at| (at, *kind)))
.min_by_key(|(at, _)| *at)
.map_or(Self::Malformed, |(_, kind)| kind)
}
fn label(self) -> &'static str {
match self {
Self::InvalidType => "invalid type",
Self::InvalidValue => "invalid value",
Self::InvalidLength => "invalid length",
Self::UnknownVariant => "unknown variant",
Self::UnknownField => "unknown field",
Self::MissingField => "missing field",
Self::DuplicateField => "duplicate field",
Self::Malformed => "malformed YAML",
Self::Serialize => "serialisation failed",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct YamlErrorDetail {
pub kind: YamlErrorKind,
pub key_path: Option<String>,
pub line: Option<usize>,
pub column: Option<usize>,
}
impl YamlErrorDetail {
fn from_parse(err: serde_path_to_error::Error<serde_yaml::Error>) -> Self {
let kind = YamlErrorKind::classify(&err.inner().to_string());
let keep = match kind {
YamlErrorKind::UnknownField => err.path().iter().len().saturating_sub(1),
_ => usize::MAX,
};
let key_path = render_key_path(err.path().iter().take(keep));
let inner = err.into_inner();
Self {
key_path,
..Self::new(kind, &inner)
}
}
fn from_serialize(err: &serde_yaml::Error) -> Self {
Self::new(YamlErrorKind::Serialize, err)
}
fn new(kind: YamlErrorKind, err: &serde_yaml::Error) -> Self {
let location = err.location();
Self {
kind,
key_path: None,
line: location.as_ref().map(serde_yaml::Location::line),
column: location.as_ref().map(serde_yaml::Location::column),
}
}
}
fn render_key_path<'a>(
segments: impl Iterator<Item = &'a serde_path_to_error::Segment>,
) -> Option<String> {
let mut out = String::new();
for segment in segments {
let is_index = matches!(segment, serde_path_to_error::Segment::Seq { .. });
if !out.is_empty() && !is_index {
out.push('.');
}
out.push_str(&segment.to_string());
}
(!out.is_empty()).then_some(out)
}
impl std::fmt::Display for YamlErrorDetail {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.kind.label())?;
if let Some(key_path) = &self.key_path {
write!(f, " at key `{key_path}`")?;
}
if let (Some(line), Some(column)) = (self.line, self.column) {
write!(f, " at line {line} column {column}")?;
}
f.write_str(" (value withheld)")
}
}
pub fn crate_config_dir_at(base: &Path, crate_name: &str) -> PathBuf {
base.join(TRUSTY_TOOLS_DIR).join(crate_name)
}
pub fn crate_config_path_at(base: &Path, crate_name: &str) -> PathBuf {
crate_config_dir_at(base, crate_name).join(CONFIG_FILE)
}
pub fn crate_config_dir(crate_name: &str) -> Option<PathBuf> {
dirs::home_dir().map(|home| crate_config_dir_at(&home, crate_name))
}
pub fn crate_config_path(crate_name: &str) -> Option<PathBuf> {
dirs::home_dir().map(|home| crate_config_path_at(&home, crate_name))
}
pub fn load_at<T: DeserializeOwned>(path: &Path) -> Result<Option<T>, ConfigError> {
let raw = match std::fs::read_to_string(path) {
Ok(raw) => raw,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(e) => {
return Err(ConfigError::Io {
path: path.to_path_buf(),
source: e,
});
}
};
let value = serde_path_to_error::deserialize::<_, T>(serde_yaml::Deserializer::from_str(&raw))
.map_err(|e| ConfigError::Yaml {
path: path.to_path_buf(),
detail: YamlErrorDetail::from_parse(e),
})?;
Ok(Some(value))
}
pub fn load<T: DeserializeOwned>(crate_name: &str) -> Result<Option<T>, ConfigError> {
match crate_config_path(crate_name) {
Some(path) => load_at(path.as_path()),
None => Ok(None),
}
}
pub fn load_or_default<T: DeserializeOwned + Default>(crate_name: &str) -> T {
match crate_config_path(crate_name) {
Some(path) => load_or_default_at(path.as_path(), crate_name),
None => T::default(),
}
}
pub fn load_or_default_at<T: DeserializeOwned + Default>(path: &Path, crate_name: &str) -> T {
match load_at::<T>(path) {
Ok(Some(value)) => value,
Ok(None) => T::default(),
Err(e) => {
tracing::warn!("{e}; falling back to default {crate_name} config");
T::default()
}
}
}
pub fn save_at<T: Serialize>(path: &Path, value: &T) -> Result<PathBuf, ConfigError> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).map_err(|e| ConfigError::Io {
path: parent.to_path_buf(),
source: e,
})?;
}
let yaml = serde_yaml::to_string(value).map_err(|e| ConfigError::Yaml {
path: path.to_path_buf(),
detail: YamlErrorDetail::from_serialize(&e), })?;
let header = "# .trusty-tools/<crate>/config.yaml\n\
# Managed by the trusty-tools config convention (#1220).\n\
# Edit by hand or via the trusty-console Config tab.\n\n";
let content = format!("{header}{yaml}");
save_raw_at(path, &content)
}
pub fn save_raw_at(path: &Path, contents: &str) -> Result<PathBuf, ConfigError> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).map_err(|e| ConfigError::Io {
path: parent.to_path_buf(),
source: e,
})?;
}
let tmp = path.with_extension("yaml.tmp");
std::fs::write(&tmp, contents).map_err(|e| ConfigError::Io {
path: tmp.clone(),
source: e,
})?;
std::fs::rename(&tmp, path).map_err(|e| ConfigError::Io {
path: path.to_path_buf(),
source: e,
})?;
Ok(path.to_path_buf())
}
pub fn save<T: Serialize>(crate_name: &str, value: &T) -> Result<PathBuf, ConfigError> {
match crate_config_path(crate_name) {
Some(path) => save_at(path.as_path(), value),
None => Err(ConfigError::Io {
path: PathBuf::from(format!("~/{TRUSTY_TOOLS_DIR}/{crate_name}/{CONFIG_FILE}")),
source: std::io::Error::new(std::io::ErrorKind::NotFound, "home directory unavailable"),
}),
}
}
#[cfg(test)]
#[path = "crate_config_redaction_tests.rs"]
mod redaction_tests;
#[cfg(test)]
mod tests {
use super::*;
use serde::Deserialize;
#[derive(Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
struct Sample {
#[serde(default)]
name: String,
#[serde(default)]
count: u32,
}
#[test]
fn crate_config_path_layout() {
let p = crate_config_path_at(Path::new("/home/bob"), "trusty-mpm");
assert_eq!(
p,
PathBuf::from("/home/bob/.trusty-tools/trusty-mpm/config.yaml")
);
let d = crate_config_dir_at(Path::new("/home/bob"), "trusty-mpm");
assert_eq!(d, PathBuf::from("/home/bob/.trusty-tools/trusty-mpm"));
}
#[test]
fn load_absent_is_none() {
let tmp = tempfile::TempDir::new().unwrap();
let path = crate_config_path_at(tmp.path(), "trusty-mpm");
let got: Option<Sample> = load_at(&path).unwrap();
assert_eq!(got, None);
}
#[test]
fn save_then_load_round_trips() {
let tmp = tempfile::TempDir::new().unwrap();
let path = crate_config_path_at(tmp.path(), "trusty-mpm");
let value = Sample {
name: "demo".into(),
count: 7,
};
let written = save_at(&path, &value).unwrap();
assert_eq!(written, path);
let got: Sample = load_at(&path).unwrap().expect("present");
assert_eq!(got, value);
let raw = std::fs::read_to_string(&path).unwrap();
assert!(raw.contains("trusty-tools config convention"));
}
#[test]
fn save_raw_at_leaves_no_tmp_sibling() {
let tmp = tempfile::TempDir::new().unwrap();
let path = crate_config_path_at(tmp.path(), "trusty-mpm");
let written = save_raw_at(&path, "name: demo\n").unwrap();
assert_eq!(written, path);
assert_eq!(std::fs::read_to_string(&path).unwrap(), "name: demo\n");
let siblings: Vec<PathBuf> = std::fs::read_dir(path.parent().unwrap())
.unwrap()
.map(|e| e.unwrap().path())
.collect();
assert_eq!(siblings, vec![path], "a .tmp sibling survived the write");
}
#[cfg(unix)]
#[test]
fn a_failed_raw_write_leaves_the_target_byte_identical() {
use std::os::unix::fs::PermissionsExt;
let tmp = tempfile::TempDir::new().unwrap();
let path = crate_config_path_at(tmp.path(), "trusty-mpm");
let dir = path.parent().unwrap().to_path_buf();
std::fs::create_dir_all(&dir).unwrap();
let original = "# hand-written\nname: operator\n";
std::fs::write(&path, original).unwrap();
let restore = std::fs::metadata(&dir).unwrap().permissions();
std::fs::set_permissions(&dir, std::fs::Permissions::from_mode(0o555)).unwrap();
let root_can_still_write = std::fs::write(dir.join("probe"), b"x").is_ok();
let result = if root_can_still_write {
std::fs::remove_file(dir.join("probe")).ok();
None
} else {
Some(save_raw_at(&path, "name: clobbered\n"))
};
std::fs::set_permissions(&dir, restore).unwrap();
let Some(result) = result else {
return; };
assert!(result.is_err(), "the write was expected to fail");
assert_eq!(
std::fs::read_to_string(&path).unwrap(),
original,
"a failed write modified the target"
);
}
#[test]
fn load_or_default_on_missing() {
let tmp = tempfile::TempDir::new().unwrap();
let path = crate_config_path_at(tmp.path(), "absent-crate");
assert_eq!(load_at::<Sample>(&path).unwrap(), None);
assert_eq!(
Sample::default(),
Sample {
name: String::new(),
count: 0
}
);
}
#[test]
fn load_malformed_is_err() {
let tmp = tempfile::TempDir::new().unwrap();
let path = crate_config_path_at(tmp.path(), "trusty-mpm");
std::fs::create_dir_all(path.parent().unwrap()).unwrap();
std::fs::write(&path, "name: demo\ncount: not-a-number\n").unwrap();
let err = load_at::<Sample>(&path).unwrap_err();
assert!(matches!(err, ConfigError::Yaml { .. }), "got {err:?}");
}
}