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_struct<T: serde::Serialize>(&self, value: &T) -> Result<(), Error> {
105 let serialized = Value::serialize(value).map_err(|error| {
106 Error::new(
107 crate::ErrorKind::Type,
108 format!("the defaults struct did not serialize: {error}"),
109 )
110 })?;
111
112 let Value::Dict(_, entries) = serialized else {
113 return Err(Error::new(
114 crate::ErrorKind::Type,
115 "defaults must be a struct or a map; a bare value has no field name to live under",
116 ));
117 };
118
119 let mut layer = self.lock();
120
121 for (path, value) in entries {
122 layer.insert(path, value);
123 }
124
125 Ok(())
126 }
127
128 pub fn set_text(&self, path: &str, text: &str) -> Result<(), Error> {
139 check_path(path)?;
140
141 let value = text
145 .parse::<Value>()
146 .unwrap_or_else(|_| Value::from(text.to_owned()));
147
148 self.lock().insert(path.to_owned(), value);
149
150 Ok(())
151 }
152
153 pub fn set_assignments<I, S>(&self, assignments: I) -> Result<(), Error>
163 where
164 I: IntoIterator<Item = S>,
165 S: AsRef<str>,
166 {
167 for assignment in assignments {
168 let assignment = assignment.as_ref();
169
170 let Some((path, value)) = assignment.split_once('=') else {
171 return Err(Error::new(
172 ErrorKind::Type,
173 format!("`{assignment}` is not a `key=value` assignment"),
174 ));
175 };
176
177 self.set_text(path.trim(), value)?;
178 }
179
180 Ok(())
181 }
182
183 #[cfg(feature = "clap")]
208 #[cfg_attr(docsrs, doc(cfg(feature = "clap")))]
209 pub fn bind_clap(
210 &self,
211 matches: &clap::ArgMatches,
212 bindings: &[(&str, &str)],
213 ) -> Result<(), Error> {
214 for (argument, path) in bindings {
215 if matches.value_source(argument) != Some(clap::parser::ValueSource::CommandLine) {
216 continue;
217 }
218
219 let Some(mut values) = matches.get_raw(argument) else {
220 continue;
221 };
222
223 let raw: Vec<&std::ffi::OsStr> = values.by_ref().collect();
226
227 let text = match raw.as_slice() {
228 [] => continue,
229 [single] => utf8(single, argument)?.to_owned(),
230 many => {
231 let mut rendered = String::from("[");
232
233 for (index, value) in many.iter().enumerate() {
234 if index > 0 {
235 rendered.push(',');
236 }
237
238 rendered.push_str(utf8(value, argument)?);
239 }
240
241 rendered.push(']');
242 rendered
243 }
244 };
245
246 self.set_text(path, &text)?;
247 }
248
249 Ok(())
250 }
251
252 #[must_use = "the return says whether anything was removed; ignore it \
254 deliberately with `let _ =` if you do not care"]
255 pub fn unset(&self, path: &str) -> bool {
256 self.lock().remove(path).is_some()
257 }
258
259 pub fn clear(&self) {
262 self.lock().clear();
263 }
264
265 #[must_use]
268 pub fn is_empty(&self) -> bool {
269 self.lock().is_empty()
270 }
271
272 fn lock(&self) -> std::sync::MutexGuard<'_, BTreeMap<String, Value>> {
280 self.entries
281 .lock()
282 .unwrap_or_else(std::sync::PoisonError::into_inner)
283 }
284
285 fn dict(&self) -> Dict {
287 let mut root = Dict::new();
288
289 for (path, value) in self.lock().iter() {
290 insert_path(&mut root, path, value.clone());
291 }
292
293 root
294 }
295
296 pub(crate) fn provider<'a>(&'a self, profile: &str, name: &'static str) -> LayerProvider<'a> {
298 LayerProvider {
299 layer: self,
300 profile: Profile::from(profile),
301 name,
302 }
303 }
304}
305
306#[cfg(feature = "clap")]
309fn utf8<'a>(value: &'a std::ffi::OsStr, argument: &str) -> Result<&'a str, Error> {
310 value.to_str().ok_or_else(|| {
311 Error::new(
312 ErrorKind::Type,
313 format!("`--{argument}` is not valid UTF-8"),
314 )
315 })
316}
317
318pub(crate) fn check_path(path: &str) -> Result<(), Error> {
327 if path.is_empty() || path.split('.').any(str::is_empty) {
328 return Err(Error::new(
329 ErrorKind::Type,
330 format!("`{path}` is not a usable key path"),
331 ));
332 }
333
334 Ok(())
335}
336
337pub(crate) fn insert_path(root: &mut Dict, path: &str, value: Value) {
338 let mut segments = path.split('.').peekable();
339 let mut current = root;
340
341 while let Some(segment) = segments.next() {
342 if segments.peek().is_none() {
343 current.insert(segment.to_owned(), value);
344 return;
345 }
346
347 let entry = current
348 .entry(segment.to_owned())
349 .or_insert_with(|| Value::from(Dict::new()));
350
351 if !matches!(entry, Value::Dict(..)) {
352 *entry = Value::from(Dict::new());
353 }
354
355 let Value::Dict(_, nested) = entry else {
356 unreachable!("just replaced with a dict")
357 };
358
359 current = nested;
360 }
361}
362
363pub(crate) struct LayerProvider<'a> {
364 layer: &'a Layer,
365 profile: Profile,
366 name: &'static str,
367}
368
369impl std::fmt::Debug for Layer {
373 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
374 let entries = self.lock();
375
376 f.debug_struct("Layer")
377 .field("keys", &entries.keys().collect::<Vec<_>>())
378 .field("len", &entries.len())
379 .finish_non_exhaustive()
380 }
381}
382
383impl Provider for LayerProvider<'_> {
384 fn metadata(&self) -> Metadata {
385 Metadata::named(self.name)
386 }
387
388 fn data(&self) -> figment::Result<Map<Profile, Dict>> {
389 let mut map = Map::new();
390 map.insert(self.profile.clone(), self.layer.dict());
391
392 Ok(map)
393 }
394}
395
396#[cfg(test)]
397mod tests {
398 use super::*;
399
400 #[test]
401 fn a_fresh_layer_contributes_nothing() {
402 assert!(Layer::new().is_empty());
403 }
404
405 #[test]
406 fn setting_the_same_path_replaces_it() {
407 let layer = Layer::new();
408 layer.set("port", 1u16).unwrap();
409 layer.set("port", 2u16).unwrap();
410
411 let dict = layer.dict();
412 assert_eq!(dict.get("port"), Some(&Value::from(2u16)));
413 }
414
415 #[test]
416 fn a_dotted_path_becomes_a_nested_table() {
417 let layer = Layer::new();
418 layer.set("pool.max_size", 32u16).unwrap();
419
420 let dict = layer.dict();
421 let Some(Value::Dict(_, pool)) = dict.get("pool") else {
422 panic!("expected a nested dict, got {dict:?}");
423 };
424
425 assert_eq!(pool.get("max_size"), Some(&Value::from(32u16)));
426 }
427
428 #[test]
429 fn siblings_under_one_parent_do_not_clobber_each_other() {
430 let layer = Layer::new();
431 layer.set("pool.max_size", 32u16).unwrap();
432 layer.set("pool.min_size", 4u16).unwrap();
433
434 let dict = layer.dict();
435 let Some(Value::Dict(_, pool)) = dict.get("pool") else {
436 panic!("expected a nested dict");
437 };
438
439 assert_eq!(pool.len(), 2);
440 }
441
442 #[test]
443 fn a_scalar_standing_where_a_table_is_needed_is_replaced() {
444 let layer = Layer::new();
445 layer.set("pool", 1u16).unwrap();
446 layer.set("pool.max_size", 32u16).unwrap();
447
448 let dict = layer.dict();
449 assert!(matches!(dict.get("pool"), Some(Value::Dict(..))));
450 }
451
452 #[test]
453 fn unset_and_clear_both_report_honestly() {
454 let layer = Layer::new();
455 layer.set("a", 1u16).unwrap();
456
457 assert!(layer.unset("a"));
458 assert!(!layer.unset("a"));
459 assert!(layer.is_empty());
460
461 layer.set("b", 1u16).unwrap();
462 layer.clear();
463 assert!(layer.is_empty());
464 }
465
466 #[test]
467 fn text_is_read_the_way_an_environment_variable_is() {
468 let layer = Layer::new();
469 layer.set_text("port", "8080").unwrap();
470 layer.set_text("enabled", "true").unwrap();
471 layer.set_text("host", "localhost").unwrap();
472
473 let dict = layer.dict();
474
475 assert_eq!(dict.get("port"), Some(&Value::from(8080u64)));
476 assert_eq!(dict.get("enabled"), Some(&Value::from(true)));
477 assert_eq!(dict.get("host"), Some(&Value::from("localhost")));
478 }
479
480 #[test]
481 fn assignments_are_split_on_the_first_equals() {
482 let layer = Layer::new();
483 layer
484 .set_assignments(["db.host=post=gres", "db.port=5432"])
485 .unwrap();
486
487 let dict = layer.dict();
488 let Some(Value::Dict(_, db)) = dict.get("db") else {
489 panic!("expected a nested dict");
490 };
491
492 assert_eq!(db.get("host"), Some(&Value::from("post=gres")));
493 assert_eq!(db.get("port"), Some(&Value::from(5432u64)));
494 }
495
496 #[test]
497 fn an_assignment_without_an_equals_names_itself() {
498 let error = Layer::new().set_assignments(["nonsense"]).unwrap_err();
499
500 assert!(error.to_string().contains("`nonsense`"), "{error}");
501 }
502
503 #[test]
504 fn an_unusable_path_is_rejected_at_the_call_site() {
505 let layer = Layer::new();
506
507 assert!(layer.set("", 1u16).is_err());
508 assert!(layer.set("a..b", 1u16).is_err());
509 assert!(layer.set(".a", 1u16).is_err());
510 }
511
512 #[test]
513 fn structured_values_survive_the_round_trip() {
514 #[derive(serde::Serialize)]
515 struct Pool {
516 max_size: u16,
517 }
518
519 let layer = Layer::new();
520 layer.set("pool", Pool { max_size: 7 }).unwrap();
521
522 let dict = layer.dict();
523 let Some(Value::Dict(_, pool)) = dict.get("pool") else {
524 panic!("expected a nested dict");
525 };
526
527 assert_eq!(pool.get("max_size"), Some(&Value::from(7u16)));
528 }
529}