Skip to main content

o_sfu/config/
env.rs

1use std::{
2    io, iter,
3    marker::PhantomData,
4    net::{IpAddr, SocketAddr},
5    num::{NonZeroU64, NonZeroUsize},
6    path::Path,
7    time::Duration,
8};
9
10use anyhow::{Context, Result, anyhow, ensure};
11use o_sfu_core::prelude::Bitrate;
12use secrecy::SecretString;
13use zeroize::Zeroize;
14
15use super::DeadlineDuration;
16
17type Lookup<'a> = dyn Fn(&str) -> Option<String> + 'a;
18type ReadFile<'a> = dyn Fn(&Path) -> io::Result<String> + 'a;
19
20enum FileSource {
21    Fallback(&'static str),
22    Exclusive(&'static str),
23}
24
25pub(super) struct EnvValue {
26    pub(super) key: &'static str,
27    pub(super) raw: String,
28}
29
30pub(super) struct Env<'a> {
31    lookup: Box<Lookup<'a>>,
32    read_file: Box<ReadFile<'a>>,
33}
34
35impl<'a> Env<'a> {
36    pub(super) fn new(
37        get_var: impl Fn(&str) -> Option<String> + 'a,
38        read_file: impl Fn(&Path) -> io::Result<String> + 'a,
39    ) -> Self {
40        Self {
41            lookup: Box::new(get_var),
42            read_file: Box::new(read_file),
43        }
44    }
45
46    pub(super) fn var<T>(&self, key: &'static str) -> Var<'a, '_, T> {
47        Var {
48            lookup: self.lookup.as_ref(),
49            read_file: self.read_file.as_ref(),
50            key,
51            check: |_key, value| Ok(value),
52            aliases: Vec::new(),
53            file_source: None,
54            value: PhantomData,
55        }
56    }
57}
58
59/// Parses the first present key and validates its value in check order.
60///
61/// Checks may transform values and capture other settings. They receive the
62/// supplying key, including aliases. Defaults use the primary key and pass
63/// through the same checks. Missing optional values bypass parsing and checks.
64pub(super) struct Var<'env, 'lookup, T, C = fn(&'static str, T) -> Result<T>> {
65    lookup: &'lookup Lookup<'env>,
66    read_file: &'lookup ReadFile<'env>,
67    key: &'static str,
68    check: C,
69    aliases: Vec<&'static str>,
70    file_source: Option<FileSource>,
71    value: PhantomData<fn(T) -> T>,
72}
73
74impl<'env, 'lookup, T, C> Var<'env, 'lookup, T, C>
75where
76    T: EnvParse,
77    C: Fn(&'static str, T) -> Result<T>,
78{
79    /// Appends a check that runs only after all preceding checks succeed.
80    pub(super) fn check(
81        self,
82        check: impl Fn(&'static str, T) -> Result<T>,
83    ) -> Var<'env, 'lookup, T, impl Fn(&'static str, T) -> Result<T>> {
84        Var {
85            lookup: self.lookup,
86            read_file: self.read_file,
87            key: self.key,
88            check: move |key, value| check(key, (self.check)(key, value)?),
89            aliases: self.aliases,
90            file_source: self.file_source,
91            value: PhantomData,
92        }
93    }
94
95    pub(super) fn alias(mut self, alias: &'static str) -> Self {
96        self.aliases.push(alias);
97        self
98    }
99
100    pub(super) fn or_load_from_file(mut self, alias: &'static str) -> Self {
101        self.file_source = Some(FileSource::Fallback(alias));
102        self
103    }
104
105    /// Loads the file only when no direct source is configured.
106    ///
107    /// Reading a direct source together with this file source returns an
108    /// `anyhow::Error` without reading the file or disclosing the direct value.
109    pub(super) fn exclusive_file(mut self, alias: &'static str) -> Self {
110        self.file_source = Some(FileSource::Exclusive(alias));
111        self
112    }
113
114    /// Returns the parsed and checked value of the first present key.
115    ///
116    /// # Errors
117    /// Returns [`anyhow::Error`] when every key is absent or parsing or a check fails.
118    pub(super) fn required(self) -> Result<T> {
119        let value = self
120            .load()?
121            .with_context(|| format!("{} env variable is required", self.key))?;
122        self.parse(value)
123    }
124
125    /// Uses the typed default only when every key is absent.
126    ///
127    /// # Errors
128    /// Returns [`anyhow::Error`] when parsing or a check fails, including checks
129    /// of the default value.
130    pub(super) fn default(self, default: T) -> Result<T> {
131        let Some(value) = self.load()? else {
132            return (self.check)(self.key, default);
133        };
134        self.parse(value)
135    }
136
137    /// Returns `None` only when every key is absent.
138    ///
139    /// # Errors
140    /// Returns [`anyhow::Error`] when parsing or a check of a present value fails.
141    pub(super) fn optional(self) -> Result<Option<T>> {
142        self.load()?.map(|value| self.parse(value)).transpose()
143    }
144
145    fn load(&self) -> Result<Option<EnvValue>> {
146        for key in iter::once(self.key).chain(self.aliases.iter().copied()) {
147            if let Some(mut raw) = (self.lookup)(key) {
148                if let Some(FileSource::Exclusive(file_key)) = self.file_source
149                    && (self.lookup)(file_key).is_some()
150                {
151                    // A conflicting secret never reaches EnvParse's zeroizing container.
152                    raw.zeroize();
153                    return Err(anyhow!("{key} conflicts with {file_key}"));
154                }
155                return Ok(Some(EnvValue { key, raw }));
156            }
157        }
158        let Some(FileSource::Fallback(file_key) | FileSource::Exclusive(file_key)) =
159            self.file_source
160        else {
161            return Ok(None);
162        };
163        match (self.lookup)(file_key) {
164            Some(path) => {
165                let mut raw = (self.read_file)(Path::new(&path)).with_context(|| {
166                    format!("{file_key} points to \"{path}\" which could not be read")
167                })?;
168                let raw_trimmed = raw.trim().to_owned();
169                raw.zeroize();
170                Ok(Some(EnvValue {
171                    key: file_key,
172                    raw: raw_trimmed,
173                }))
174            }
175            None => Ok(None),
176        }
177    }
178
179    fn parse(&self, value: EnvValue) -> Result<T> {
180        let key = value.key;
181        (self.check)(key, T::parse(value)?)
182    }
183}
184
185pub(super) trait EnvParse: Sized {
186    fn parse(value: EnvValue) -> Result<Self>;
187}
188
189macro_rules! parse_from_str {
190    ($type:ty, $name:literal) => {
191        impl EnvParse for $type {
192            fn parse(value: EnvValue) -> Result<Self> {
193                let key = value.key;
194                value
195                    .raw
196                    .parse()
197                    .map_err(|_error| anyhow!("{key} must be a valid {}", $name))
198            }
199        }
200    };
201}
202
203parse_from_str!(IpAddr, "IP address");
204parse_from_str!(SocketAddr, "socket address");
205parse_from_str!(u8, "u8");
206parse_from_str!(u16, "u16");
207parse_from_str!(u64, "u64");
208parse_from_str!(usize, "usize");
209
210impl EnvParse for NonZeroUsize {
211    fn parse(value: EnvValue) -> Result<Self> {
212        let key = value.key;
213        Self::new(usize::parse(value)?).ok_or_else(|| anyhow!("{key} must be greater than zero"))
214    }
215}
216
217impl EnvParse for NonZeroU64 {
218    fn parse(value: EnvValue) -> Result<Self> {
219        let key = value.key;
220        Self::new(u64::parse(value)?).ok_or_else(|| anyhow!("{key} must be greater than zero"))
221    }
222}
223
224/// Parses integer bits per second, including zero.
225impl EnvParse for Bitrate {
226    fn parse(value: EnvValue) -> Result<Self> {
227        u64::parse(value).map(Self::from_bps)
228    }
229}
230
231impl EnvParse for bool {
232    fn parse(value: EnvValue) -> Result<Self> {
233        let key = value.key;
234        value
235            .raw
236            .parse()
237            .map_err(|_error| anyhow!("{key} must be either `true` or `false`"))
238    }
239}
240
241impl EnvParse for String {
242    fn parse(value: EnvValue) -> Result<Self> {
243        Ok(value.raw)
244    }
245}
246
247impl EnvParse for DeadlineDuration {
248    fn parse(value: EnvValue) -> Result<Self> {
249        let key = value.key;
250        Self::from_millis(u64::parse(value)?).map_err(|error| anyhow!("{key} {error}"))
251    }
252}
253
254impl EnvParse for Duration {
255    fn parse(value: EnvValue) -> Result<Self> {
256        let key = value.key;
257        let seconds = value
258            .raw
259            .parse()
260            .map_err(|_error| anyhow!("{key} must be a valid duration in seconds"))?;
261        Ok(Self::from_secs(seconds))
262    }
263}
264
265impl EnvParse for SecretString {
266    fn parse(value: EnvValue) -> Result<Self> {
267        Ok(Self::from(value.raw))
268    }
269}
270
271/// Requires a value greater than zero for types whose default is zero.
272///
273/// # Errors
274/// Returns [`anyhow::Error`] when the value is not greater than zero.
275pub(super) fn positive<T>(key: &'static str, value: T) -> Result<T>
276where
277    T: Default + PartialOrd,
278{
279    ensure!(value > T::default(), "{key} must be greater than zero");
280    Ok(value)
281}
282
283pub(super) fn non_empty(key: &'static str, value: String) -> Result<String> {
284    let trimmed = value.trim();
285    ensure!(!trimmed.is_empty(), "{key} must not be empty");
286    if trimmed.len() == value.len() {
287        Ok(value)
288    } else {
289        Ok(trimmed.to_owned())
290    }
291}
292
293#[cfg(test)]
294#[path = "TESTS/env.rs"]
295mod tests;