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
59pub(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 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 pub(super) fn exclusive_file(mut self, alias: &'static str) -> Self {
110 self.file_source = Some(FileSource::Exclusive(alias));
111 self
112 }
113
114 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 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 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 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
224impl 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
271pub(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;