use std::path::{Path, PathBuf};
use serde::de::DeserializeOwned;
use crate::error::{ConfigError, locate};
pub trait AppConfig: DeserializeOwned + Default {
const APPLICATION: &'static str;
const FILE: &'static str = "config.toml";
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ConfigSource {
File(PathBuf),
Defaults,
}
fn map_xdg(e: hjkl_xdg::Error) -> ConfigError {
match e {
hjkl_xdg::Error::NoHomeDir => ConfigError::NoHomeDir,
}
}
pub fn config_dir(app: &str) -> Result<PathBuf, ConfigError> {
hjkl_xdg::config_dir(app).map_err(map_xdg)
}
pub fn data_dir(app: &str) -> Result<PathBuf, ConfigError> {
hjkl_xdg::data_dir(app).map_err(map_xdg)
}
pub fn cache_dir(app: &str) -> Result<PathBuf, ConfigError> {
hjkl_xdg::cache_dir(app).map_err(map_xdg)
}
pub fn config_path<C: AppConfig>() -> Result<PathBuf, ConfigError> {
Ok(config_dir(C::APPLICATION)?.join(C::FILE))
}
pub fn load<C: AppConfig>() -> Result<(C, ConfigSource), ConfigError> {
let path = match config_path::<C>() {
Ok(p) => p,
Err(ConfigError::NoHomeDir) => {
return Ok((C::default(), ConfigSource::Defaults));
}
Err(e) => return Err(e),
};
if !path.exists() {
return Ok((C::default(), ConfigSource::Defaults));
}
let cfg = load_from::<C>(&path)?;
Ok((cfg, ConfigSource::File(path)))
}
pub fn load_from<C: AppConfig>(path: &Path) -> Result<C, ConfigError> {
let src = std::fs::read_to_string(path).map_err(|e| ConfigError::Io {
path: path.to_path_buf(),
source: e,
})?;
toml::from_str::<C>(&src).map_err(|e| {
let span = e.span().unwrap_or(0..0);
let (line, col, snippet) = locate(&src, span.start);
ConfigError::Parse {
path: path.to_path_buf(),
line,
col,
message: e.message().to_string(),
snippet,
}
})
}
pub fn load_layered<C: AppConfig>(defaults_toml: &str) -> Result<(C, ConfigSource), ConfigError> {
let user_path = match config_path::<C>() {
Ok(p) if p.exists() => Some(p),
Ok(_) => None,
Err(ConfigError::NoHomeDir) => None,
Err(e) => return Err(e),
};
match user_path {
Some(p) => {
let cfg = load_layered_from::<C>(defaults_toml, &p)?;
Ok((cfg, ConfigSource::File(p)))
}
None => {
let cfg = parse_defaults_only::<C>(defaults_toml)?;
Ok((cfg, ConfigSource::Defaults))
}
}
}
pub fn load_layered_from<C: AppConfig>(defaults_toml: &str, path: &Path) -> Result<C, ConfigError> {
let user_src = std::fs::read_to_string(path).map_err(|e| ConfigError::Io {
path: path.to_path_buf(),
source: e,
})?;
let user_table: toml::Table = toml::from_str(&user_src).map_err(|e| {
let span = e.span().unwrap_or(0..0);
let (line, col, snippet) = locate(&user_src, span.start);
ConfigError::Parse {
path: path.to_path_buf(),
line,
col,
message: e.message().to_string(),
snippet,
}
})?;
let mut merged: toml::Table =
toml::from_str(defaults_toml).map_err(|e| ConfigError::Invalid {
path: PathBuf::from("<bundled defaults>"),
message: format!("bundled defaults TOML is invalid: {e}"),
})?;
deep_merge(&mut merged, user_table);
toml::Value::Table(merged)
.try_into()
.map_err(|e: toml::de::Error| ConfigError::Invalid {
path: path.to_path_buf(),
message: e.to_string(),
})
}
fn parse_defaults_only<C: AppConfig>(defaults_toml: &str) -> Result<C, ConfigError> {
toml::from_str::<C>(defaults_toml).map_err(|e| ConfigError::Invalid {
path: PathBuf::from("<bundled defaults>"),
message: format!("bundled defaults TOML is invalid: {e}"),
})
}
pub fn deep_merge(into: &mut toml::Table, from: toml::Table) {
for (k, v) in from {
match (into.get_mut(&k), v) {
(Some(toml::Value::Table(into_t)), toml::Value::Table(from_t)) => {
deep_merge(into_t, from_t);
}
(_, v) => {
into.insert(k, v);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dirs_delegate_to_hjkl_xdg() {
assert_eq!(
config_dir("myapp").unwrap(),
hjkl_xdg::config_home().unwrap().join("myapp")
);
assert_eq!(
data_dir("myapp").unwrap(),
hjkl_xdg::data_home().unwrap().join("myapp")
);
assert_eq!(
cache_dir("myapp").unwrap(),
hjkl_xdg::cache_home().unwrap().join("myapp")
);
}
#[test]
fn xdg_no_home_maps_to_no_home_dir() {
assert!(matches!(
map_xdg(hjkl_xdg::Error::NoHomeDir),
ConfigError::NoHomeDir
));
}
#[test]
fn config_dir_smoke() {
let p = config_dir("myapp").unwrap();
assert!(
p.ends_with("myapp"),
"expected path ending in `myapp`, got {p:?}"
);
}
#[test]
fn data_dir_smoke() {
let p = data_dir("myapp").unwrap();
assert!(p.ends_with("myapp"));
}
#[test]
fn cache_dir_smoke() {
let p = cache_dir("myapp").unwrap();
assert!(p.ends_with("myapp"));
}
}