Skip to main content

rama_http/layer/dns/dns_resolve/
username_parser.rs

1use super::DnsResolveMode;
2use rama_core::error::BoxErrorExt as _;
3use rama_core::username::{UsernameLabelParser, UsernameLabelState};
4use rama_core::{
5    error::{BoxError, ErrorContext},
6    extensions::Extensions,
7    telemetry::tracing,
8};
9use rama_utils::macros::str::eq_ignore_ascii_case;
10
11#[derive(Debug, Clone, Default)]
12#[non_exhaustive]
13/// A parser which parses [`DnsResolveMode`]s from username labels
14/// and adds it to the [`Extensions`].
15///
16/// [`Extensions`]: rama_core::extensions::Extensions
17pub struct DnsResolveModeUsernameParser {
18    key_found: bool,
19    mode: DnsResolveMode,
20}
21
22impl DnsResolveModeUsernameParser {
23    /// Create a new [`DnsResolveModeUsernameParser`].
24    #[must_use]
25    pub fn new() -> Self {
26        Self::default()
27    }
28}
29
30impl UsernameLabelParser for DnsResolveModeUsernameParser {
31    type Error = BoxError;
32
33    fn parse_label(&mut self, label: &str) -> UsernameLabelState {
34        if self.key_found {
35            self.mode = match label
36                .parse()
37                .context("parse dns resolve mode username label")
38            {
39                Ok(mode) => mode,
40                Err(err) => {
41                    tracing::trace!("abort username label parsing: invalid parse label: {err:?}");
42                    return UsernameLabelState::Abort;
43                }
44            };
45            self.key_found = false;
46            UsernameLabelState::Used
47        } else if eq_ignore_ascii_case!("dns", label) {
48            self.key_found = true;
49            UsernameLabelState::Used
50        } else {
51            UsernameLabelState::Ignored
52        }
53    }
54
55    fn build(self, ext: &Extensions) -> Result<(), Self::Error> {
56        if self.key_found {
57            return Err(BoxError::from_static_str(
58                "unused dns resolve mode username key: dns",
59            ));
60        }
61        ext.insert(self.mode);
62        Ok(())
63    }
64}
65
66#[cfg(test)]
67mod tests {
68    use super::*;
69    use rama_core::username::parse_username;
70
71    #[test]
72    fn test_username_dns_resolve_mod_config() {
73        let test_cases = [
74            ("john", String::from("john"), DnsResolveMode::default()),
75            (
76                "john-dns-eager",
77                String::from("john"),
78                DnsResolveMode::eager(),
79            ),
80            (
81                "john-dns-lazy",
82                String::from("john"),
83                DnsResolveMode::lazy(),
84            ),
85            (
86                "john-dns-eager-dns-lazy",
87                String::from("john"),
88                DnsResolveMode::lazy(),
89            ),
90            (
91                "john-dns-lazy-dns-eager",
92                String::from("john"),
93                DnsResolveMode::eager(),
94            ),
95        ];
96
97        for (username, expected_username, expected_mode) in test_cases.into_iter() {
98            let ext = Extensions::default();
99
100            let parser = DnsResolveModeUsernameParser::default();
101
102            let username = parse_username(&ext, parser, username).unwrap();
103            let mode = *ext.get_ref::<DnsResolveMode>().unwrap();
104            assert_eq!(
105                username, expected_username,
106                "username = '{username}' ; expected_username = '{expected_username}'",
107            );
108            assert_eq!(
109                mode, expected_mode,
110                "username = '{username}' ; expected_mode = '{expected_mode}'",
111            );
112        }
113    }
114
115    #[test]
116    fn test_username_dns_resolve_mode_error() {
117        for username in [
118            "john-",
119            "john-dns",
120            "john-dns-eager-",
121            "john-dns-eager-dns",
122            "john-dns-foo",
123        ] {
124            let ext = Extensions::default();
125
126            let parser = DnsResolveModeUsernameParser::default();
127
128            assert!(
129                parse_username(&ext, parser, username).is_err(),
130                "username = {username}",
131            );
132        }
133    }
134}