1use std::collections::BTreeMap;
20use std::sync::Mutex;
21
22use figment::value::{Dict, Map, Value};
23use figment::{Metadata, Profile, Provider};
24use serde::Serialize;
25
26use crate::error::{Error, ErrorKind};
27
28pub(crate) const DEFAULTS_NAME: &str = "values set as defaults";
30
31pub(crate) const OVERRIDES_NAME: &str = "values set as overrides";
33
34pub(crate) const FLAGS_NAME: &str = "values set from the command line";
36
37#[derive(Default)]
58pub struct Layer {
59 entries: Mutex<BTreeMap<String, Value>>,
60}
61
62impl Layer {
63 #[must_use]
65 pub const fn new() -> Self {
66 Self {
67 entries: Mutex::new(BTreeMap::new()),
68 }
69 }
70
71 pub fn set<T: Serialize>(&self, path: &str, value: T) -> Result<(), Error> {
83 check_path(path)?;
84
85 let value = Value::serialize(value)
86 .map_err(|error| Error::new(ErrorKind::Type, error.to_string()).prepend_key(path))?;
87
88 self.lock().insert(path.to_owned(), value);
89
90 Ok(())
91 }
92
93 pub fn set_text(&self, path: &str, text: &str) -> Result<(), Error> {
104 check_path(path)?;
105
106 let value = text
110 .parse::<Value>()
111 .unwrap_or_else(|_| Value::from(text.to_owned()));
112
113 self.lock().insert(path.to_owned(), value);
114
115 Ok(())
116 }
117
118 pub fn set_assignments<I, S>(&self, assignments: I) -> Result<(), Error>
128 where
129 I: IntoIterator<Item = S>,
130 S: AsRef<str>,
131 {
132 for assignment in assignments {
133 let assignment = assignment.as_ref();
134
135 let Some((path, value)) = assignment.split_once('=') else {
136 return Err(Error::new(
137 ErrorKind::Type,
138 format!("`{assignment}` is not a `key=value` assignment"),
139 ));
140 };
141
142 self.set_text(path.trim(), value)?;
143 }
144
145 Ok(())
146 }
147
148 #[cfg(feature = "clap")]
173 #[cfg_attr(docsrs, doc(cfg(feature = "clap")))]
174 pub fn bind_clap(
175 &self,
176 matches: &clap::ArgMatches,
177 bindings: &[(&str, &str)],
178 ) -> Result<(), Error> {
179 for (argument, path) in bindings {
180 if matches.value_source(argument) != Some(clap::parser::ValueSource::CommandLine) {
181 continue;
182 }
183
184 let Some(mut values) = matches.get_raw(argument) else {
185 continue;
186 };
187
188 let raw: Vec<&std::ffi::OsStr> = values.by_ref().collect();
191
192 let text = match raw.as_slice() {
193 [] => continue,
194 [single] => utf8(single, argument)?.to_owned(),
195 many => {
196 let mut rendered = String::from("[");
197
198 for (index, value) in many.iter().enumerate() {
199 if index > 0 {
200 rendered.push(',');
201 }
202
203 rendered.push_str(utf8(value, argument)?);
204 }
205
206 rendered.push(']');
207 rendered
208 }
209 };
210
211 self.set_text(path, &text)?;
212 }
213
214 Ok(())
215 }
216
217 pub fn unset(&self, path: &str) -> bool {
220 self.lock().remove(path).is_some()
221 }
222
223 pub fn clear(&self) {
226 self.lock().clear();
227 }
228
229 pub fn is_empty(&self) -> bool {
232 self.lock().is_empty()
233 }
234
235 fn lock(&self) -> std::sync::MutexGuard<'_, BTreeMap<String, Value>> {
243 self.entries
244 .lock()
245 .unwrap_or_else(std::sync::PoisonError::into_inner)
246 }
247
248 fn dict(&self) -> Dict {
250 let mut root = Dict::new();
251
252 for (path, value) in self.lock().iter() {
253 insert_path(&mut root, path, value.clone());
254 }
255
256 root
257 }
258
259 pub(crate) fn provider<'a>(&'a self, profile: &str, name: &'static str) -> LayerProvider<'a> {
261 LayerProvider {
262 layer: self,
263 profile: Profile::from(profile),
264 name,
265 }
266 }
267}
268
269#[cfg(feature = "clap")]
272fn utf8<'a>(value: &'a std::ffi::OsStr, argument: &str) -> Result<&'a str, Error> {
273 value.to_str().ok_or_else(|| {
274 Error::new(
275 ErrorKind::Type,
276 format!("`--{argument}` is not valid UTF-8"),
277 )
278 })
279}
280
281pub(crate) fn check_path(path: &str) -> Result<(), Error> {
290 if path.is_empty() || path.split('.').any(str::is_empty) {
291 return Err(Error::new(
292 ErrorKind::Type,
293 format!("`{path}` is not a usable key path"),
294 ));
295 }
296
297 Ok(())
298}
299
300pub(crate) fn insert_path(root: &mut Dict, path: &str, value: Value) {
301 let mut segments = path.split('.').peekable();
302 let mut current = root;
303
304 while let Some(segment) = segments.next() {
305 if segments.peek().is_none() {
306 current.insert(segment.to_owned(), value);
307 return;
308 }
309
310 let entry = current
311 .entry(segment.to_owned())
312 .or_insert_with(|| Value::from(Dict::new()));
313
314 if !matches!(entry, Value::Dict(..)) {
315 *entry = Value::from(Dict::new());
316 }
317
318 let Value::Dict(_, nested) = entry else {
319 unreachable!("just replaced with a dict")
320 };
321
322 current = nested;
323 }
324}
325
326pub(crate) struct LayerProvider<'a> {
327 layer: &'a Layer,
328 profile: Profile,
329 name: &'static str,
330}
331
332impl Provider for LayerProvider<'_> {
333 fn metadata(&self) -> Metadata {
334 Metadata::named(self.name)
335 }
336
337 fn data(&self) -> figment::Result<Map<Profile, Dict>> {
338 let mut map = Map::new();
339 map.insert(self.profile.clone(), self.layer.dict());
340
341 Ok(map)
342 }
343}
344
345#[cfg(test)]
346mod tests {
347 use super::*;
348
349 #[test]
350 fn a_fresh_layer_contributes_nothing() {
351 assert!(Layer::new().is_empty());
352 }
353
354 #[test]
355 fn setting_the_same_path_replaces_it() {
356 let layer = Layer::new();
357 layer.set("port", 1u16).unwrap();
358 layer.set("port", 2u16).unwrap();
359
360 let dict = layer.dict();
361 assert_eq!(dict.get("port"), Some(&Value::from(2u16)));
362 }
363
364 #[test]
365 fn a_dotted_path_becomes_a_nested_table() {
366 let layer = Layer::new();
367 layer.set("pool.max_size", 32u16).unwrap();
368
369 let dict = layer.dict();
370 let Some(Value::Dict(_, pool)) = dict.get("pool") else {
371 panic!("expected a nested dict, got {dict:?}");
372 };
373
374 assert_eq!(pool.get("max_size"), Some(&Value::from(32u16)));
375 }
376
377 #[test]
378 fn siblings_under_one_parent_do_not_clobber_each_other() {
379 let layer = Layer::new();
380 layer.set("pool.max_size", 32u16).unwrap();
381 layer.set("pool.min_size", 4u16).unwrap();
382
383 let dict = layer.dict();
384 let Some(Value::Dict(_, pool)) = dict.get("pool") else {
385 panic!("expected a nested dict");
386 };
387
388 assert_eq!(pool.len(), 2);
389 }
390
391 #[test]
392 fn a_scalar_standing_where_a_table_is_needed_is_replaced() {
393 let layer = Layer::new();
394 layer.set("pool", 1u16).unwrap();
395 layer.set("pool.max_size", 32u16).unwrap();
396
397 let dict = layer.dict();
398 assert!(matches!(dict.get("pool"), Some(Value::Dict(..))));
399 }
400
401 #[test]
402 fn unset_and_clear_both_report_honestly() {
403 let layer = Layer::new();
404 layer.set("a", 1u16).unwrap();
405
406 assert!(layer.unset("a"));
407 assert!(!layer.unset("a"));
408 assert!(layer.is_empty());
409
410 layer.set("b", 1u16).unwrap();
411 layer.clear();
412 assert!(layer.is_empty());
413 }
414
415 #[test]
416 fn text_is_read_the_way_an_environment_variable_is() {
417 let layer = Layer::new();
418 layer.set_text("port", "8080").unwrap();
419 layer.set_text("enabled", "true").unwrap();
420 layer.set_text("host", "localhost").unwrap();
421
422 let dict = layer.dict();
423
424 assert_eq!(dict.get("port"), Some(&Value::from(8080u64)));
425 assert_eq!(dict.get("enabled"), Some(&Value::from(true)));
426 assert_eq!(dict.get("host"), Some(&Value::from("localhost")));
427 }
428
429 #[test]
430 fn assignments_are_split_on_the_first_equals() {
431 let layer = Layer::new();
432 layer
433 .set_assignments(["db.host=post=gres", "db.port=5432"])
434 .unwrap();
435
436 let dict = layer.dict();
437 let Some(Value::Dict(_, db)) = dict.get("db") else {
438 panic!("expected a nested dict");
439 };
440
441 assert_eq!(db.get("host"), Some(&Value::from("post=gres")));
442 assert_eq!(db.get("port"), Some(&Value::from(5432u64)));
443 }
444
445 #[test]
446 fn an_assignment_without_an_equals_names_itself() {
447 let error = Layer::new().set_assignments(["nonsense"]).unwrap_err();
448
449 assert!(error.to_string().contains("`nonsense`"), "{error}");
450 }
451
452 #[test]
453 fn an_unusable_path_is_rejected_at_the_call_site() {
454 let layer = Layer::new();
455
456 assert!(layer.set("", 1u16).is_err());
457 assert!(layer.set("a..b", 1u16).is_err());
458 assert!(layer.set(".a", 1u16).is_err());
459 }
460
461 #[test]
462 fn structured_values_survive_the_round_trip() {
463 #[derive(serde::Serialize)]
464 struct Pool {
465 max_size: u16,
466 }
467
468 let layer = Layer::new();
469 layer.set("pool", Pool { max_size: 7 }).unwrap();
470
471 let dict = layer.dict();
472 let Some(Value::Dict(_, pool)) = dict.get("pool") else {
473 panic!("expected a nested dict");
474 };
475
476 assert_eq!(pool.get("max_size"), Some(&Value::from(7u16)));
477 }
478}