use std::collections::BTreeMap;
use std::sync::Mutex;
use figment::value::{Dict, Map, Value};
use figment::{Metadata, Profile, Provider};
use serde::Serialize;
use crate::error::{Error, ErrorKind};
pub(crate) const DEFAULTS_NAME: &str = "values set as defaults";
pub(crate) const OVERRIDES_NAME: &str = "values set as overrides";
pub(crate) const FLAGS_NAME: &str = "values set from the command line";
#[derive(Default)]
pub struct Layer {
entries: Mutex<BTreeMap<String, Value>>,
}
impl Layer {
#[must_use]
pub const fn new() -> Self {
Self {
entries: Mutex::new(BTreeMap::new()),
}
}
pub fn set<T: Serialize>(&self, path: &str, value: T) -> Result<(), Error> {
check_path(path)?;
let value = Value::serialize(value)
.map_err(|error| Error::new(ErrorKind::Type, error.to_string()).prepend_key(path))?;
self.lock().insert(path.to_owned(), value);
Ok(())
}
pub fn set_struct<T: serde::Serialize>(&self, value: &T) -> Result<(), Error> {
let serialized = Value::serialize(value).map_err(|error| {
Error::new(
crate::ErrorKind::Type,
format!("the defaults struct did not serialize: {error}"),
)
})?;
let Value::Dict(_, entries) = serialized else {
return Err(Error::new(
crate::ErrorKind::Type,
"defaults must be a struct or a map; a bare value has no field name to live under",
));
};
let mut layer = self.lock();
for (path, value) in entries {
layer.insert(path, value);
}
Ok(())
}
pub fn set_text(&self, path: &str, text: &str) -> Result<(), Error> {
check_path(path)?;
let value = text
.parse::<Value>()
.unwrap_or_else(|_| Value::from(text.to_owned()));
self.lock().insert(path.to_owned(), value);
Ok(())
}
pub fn set_assignments<I, S>(&self, assignments: I) -> Result<(), Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
for assignment in assignments {
let assignment = assignment.as_ref();
let Some((path, value)) = assignment.split_once('=') else {
return Err(Error::new(
ErrorKind::Type,
format!("`{assignment}` is not a `key=value` assignment"),
));
};
self.set_text(path.trim(), value)?;
}
Ok(())
}
#[cfg(feature = "clap")]
#[cfg_attr(docsrs, doc(cfg(feature = "clap")))]
pub fn bind_clap(
&self,
matches: &clap::ArgMatches,
bindings: &[(&str, &str)],
) -> Result<(), Error> {
for (argument, path) in bindings {
if matches.value_source(argument) != Some(clap::parser::ValueSource::CommandLine) {
continue;
}
let Some(mut values) = matches.get_raw(argument) else {
continue;
};
let raw: Vec<&std::ffi::OsStr> = values.by_ref().collect();
let text = match raw.as_slice() {
[] => continue,
[single] => utf8(single, argument)?.to_owned(),
many => {
let mut rendered = String::from("[");
for (index, value) in many.iter().enumerate() {
if index > 0 {
rendered.push(',');
}
rendered.push_str(utf8(value, argument)?);
}
rendered.push(']');
rendered
}
};
self.set_text(path, &text)?;
}
Ok(())
}
#[must_use = "the return says whether anything was removed; ignore it \
deliberately with `let _ =` if you do not care"]
pub fn unset(&self, path: &str) -> bool {
self.lock().remove(path).is_some()
}
pub fn clear(&self) {
self.lock().clear();
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.lock().is_empty()
}
fn lock(&self) -> std::sync::MutexGuard<'_, BTreeMap<String, Value>> {
self.entries
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn dict(&self) -> Dict {
let mut root = Dict::new();
for (path, value) in self.lock().iter() {
insert_path(&mut root, path, value.clone());
}
root
}
pub(crate) fn provider<'a>(&'a self, profile: &str, name: &'static str) -> LayerProvider<'a> {
LayerProvider {
layer: self,
profile: Profile::from(profile),
name,
}
}
}
#[cfg(feature = "clap")]
fn utf8<'a>(value: &'a std::ffi::OsStr, argument: &str) -> Result<&'a str, Error> {
value.to_str().ok_or_else(|| {
Error::new(
ErrorKind::Type,
format!("`--{argument}` is not valid UTF-8"),
)
})
}
pub(crate) fn check_path(path: &str) -> Result<(), Error> {
if path.is_empty() || path.split('.').any(str::is_empty) {
return Err(Error::new(
ErrorKind::Type,
format!("`{path}` is not a usable key path"),
));
}
Ok(())
}
pub(crate) fn insert_path(root: &mut Dict, path: &str, value: Value) {
let mut segments = path.split('.').peekable();
let mut current = root;
while let Some(segment) = segments.next() {
if segments.peek().is_none() {
current.insert(segment.to_owned(), value);
return;
}
let entry = current
.entry(segment.to_owned())
.or_insert_with(|| Value::from(Dict::new()));
if !matches!(entry, Value::Dict(..)) {
*entry = Value::from(Dict::new());
}
let Value::Dict(_, nested) = entry else {
unreachable!("just replaced with a dict")
};
current = nested;
}
}
pub(crate) struct LayerProvider<'a> {
layer: &'a Layer,
profile: Profile,
name: &'static str,
}
impl std::fmt::Debug for Layer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let entries = self.lock();
f.debug_struct("Layer")
.field("keys", &entries.keys().collect::<Vec<_>>())
.field("len", &entries.len())
.finish_non_exhaustive()
}
}
impl Provider for LayerProvider<'_> {
fn metadata(&self) -> Metadata {
Metadata::named(self.name)
}
fn data(&self) -> figment::Result<Map<Profile, Dict>> {
let mut map = Map::new();
map.insert(self.profile.clone(), self.layer.dict());
Ok(map)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_fresh_layer_contributes_nothing() {
assert!(Layer::new().is_empty());
}
#[test]
fn setting_the_same_path_replaces_it() {
let layer = Layer::new();
layer.set("port", 1u16).unwrap();
layer.set("port", 2u16).unwrap();
let dict = layer.dict();
assert_eq!(dict.get("port"), Some(&Value::from(2u16)));
}
#[test]
fn a_dotted_path_becomes_a_nested_table() {
let layer = Layer::new();
layer.set("pool.max_size", 32u16).unwrap();
let dict = layer.dict();
let Some(Value::Dict(_, pool)) = dict.get("pool") else {
panic!("expected a nested dict, got {dict:?}");
};
assert_eq!(pool.get("max_size"), Some(&Value::from(32u16)));
}
#[test]
fn siblings_under_one_parent_do_not_clobber_each_other() {
let layer = Layer::new();
layer.set("pool.max_size", 32u16).unwrap();
layer.set("pool.min_size", 4u16).unwrap();
let dict = layer.dict();
let Some(Value::Dict(_, pool)) = dict.get("pool") else {
panic!("expected a nested dict");
};
assert_eq!(pool.len(), 2);
}
#[test]
fn a_scalar_standing_where_a_table_is_needed_is_replaced() {
let layer = Layer::new();
layer.set("pool", 1u16).unwrap();
layer.set("pool.max_size", 32u16).unwrap();
let dict = layer.dict();
assert!(matches!(dict.get("pool"), Some(Value::Dict(..))));
}
#[test]
fn unset_and_clear_both_report_honestly() {
let layer = Layer::new();
layer.set("a", 1u16).unwrap();
assert!(layer.unset("a"));
assert!(!layer.unset("a"));
assert!(layer.is_empty());
layer.set("b", 1u16).unwrap();
layer.clear();
assert!(layer.is_empty());
}
#[test]
fn text_is_read_the_way_an_environment_variable_is() {
let layer = Layer::new();
layer.set_text("port", "8080").unwrap();
layer.set_text("enabled", "true").unwrap();
layer.set_text("host", "localhost").unwrap();
let dict = layer.dict();
assert_eq!(dict.get("port"), Some(&Value::from(8080u64)));
assert_eq!(dict.get("enabled"), Some(&Value::from(true)));
assert_eq!(dict.get("host"), Some(&Value::from("localhost")));
}
#[test]
fn assignments_are_split_on_the_first_equals() {
let layer = Layer::new();
layer
.set_assignments(["db.host=post=gres", "db.port=5432"])
.unwrap();
let dict = layer.dict();
let Some(Value::Dict(_, db)) = dict.get("db") else {
panic!("expected a nested dict");
};
assert_eq!(db.get("host"), Some(&Value::from("post=gres")));
assert_eq!(db.get("port"), Some(&Value::from(5432u64)));
}
#[test]
fn an_assignment_without_an_equals_names_itself() {
let error = Layer::new().set_assignments(["nonsense"]).unwrap_err();
assert!(error.to_string().contains("`nonsense`"), "{error}");
}
#[test]
fn an_unusable_path_is_rejected_at_the_call_site() {
let layer = Layer::new();
assert!(layer.set("", 1u16).is_err());
assert!(layer.set("a..b", 1u16).is_err());
assert!(layer.set(".a", 1u16).is_err());
}
#[test]
fn structured_values_survive_the_round_trip() {
#[derive(serde::Serialize)]
struct Pool {
max_size: u16,
}
let layer = Layer::new();
layer.set("pool", Pool { max_size: 7 }).unwrap();
let dict = layer.dict();
let Some(Value::Dict(_, pool)) = dict.get("pool") else {
panic!("expected a nested dict");
};
assert_eq!(pool.get("max_size"), Some(&Value::from(7u16)));
}
}