Skip to main content

mail_auth/spf/
macros.rs

1/*
2 * SPDX-FileCopyrightText: 2020 Stalwart Labs LLC <hello@stalw.art>
3 *
4 * SPDX-License-Identifier: Apache-2.0 OR MIT
5 */
6
7use super::{Macro, Variable, Variables};
8use crate::SystemTime;
9use crate::common::resolver::{decimal_u8, hex_nibble};
10use std::{borrow::Cow, net::IpAddr};
11
12const DEFAULT_DELIMITERS: u64 = 1u64 << (b'.' - b'+');
13const LAST_DELIMITER: u8 = b'_' - b'+';
14
15impl Macro {
16    pub fn eval<'z, 'x: 'z>(
17        &'z self,
18        vars: &'x Variables<'x>,
19        default: &'x str,
20        fqdn: bool,
21    ) -> Cow<'z, str> {
22        match self {
23            Macro::Literal(literal) => std::str::from_utf8(literal).unwrap_or_default().into(),
24            Macro::Variable {
25                letter,
26                num_parts,
27                reverse,
28                escape,
29                delimiters,
30            } => match vars.get(*letter, *num_parts, *reverse, *escape, fqdn, *delimiters) {
31                Cow::Borrowed(bytes) => std::str::from_utf8(bytes).unwrap_or_default().into(),
32                Cow::Owned(bytes) => String::from_utf8(bytes).unwrap_or_default().into(),
33            },
34            Macro::List(list) => {
35                let mut result = Vec::with_capacity(
36                    list.iter()
37                        .map(|item| match item {
38                            Macro::Literal(literal) => literal.len(),
39                            _ => 24,
40                        })
41                        .sum::<usize>()
42                        + 1,
43                );
44                for item in list {
45                    match item {
46                        Macro::Literal(literal) => {
47                            result.extend_from_slice(literal);
48                        }
49                        Macro::Variable {
50                            letter,
51                            num_parts,
52                            reverse,
53                            escape,
54                            delimiters,
55                        } => {
56                            vars.append(
57                                &mut result,
58                                *letter,
59                                Transform {
60                                    num_parts: *num_parts,
61                                    reverse: *reverse,
62                                    escape: *escape,
63                                    fqdn: false,
64                                    delimiters: *delimiters,
65                                },
66                            );
67                        }
68                        Macro::List(_) | Macro::None => unreachable!(),
69                    }
70                }
71                if fqdn && matches!(result.last(), Some(last) if *last != b'.') {
72                    result.push(b'.');
73                }
74                String::from_utf8(result).unwrap_or_default().into()
75            }
76            Macro::None => default.into(),
77        }
78    }
79
80    pub fn needs_ptr(&self) -> bool {
81        match self {
82            Macro::Variable { letter, .. } => *letter == Variable::ValidatedDomain,
83            Macro::List(list) => list.iter().any(|m| matches!(m, Macro::Variable { letter, .. } if *letter == Variable::ValidatedDomain)),
84            _ => false,
85        }
86    }
87}
88
89impl<'x> Variables<'x> {
90    pub fn new() -> Self {
91        Variables {
92            current_time_on_demand: true,
93            ..Default::default()
94        }
95    }
96
97    pub fn set_ip(&mut self, value: &IpAddr) {
98        let (v, i, c): (&'static [u8], Vec<u8>, Vec<u8>) = match value {
99            IpAddr::V4(ip) => {
100                let mut dotted = Vec::with_capacity(15);
101                let mut buf = [0u8; 3];
102                for octet in ip.octets() {
103                    if !dotted.is_empty() {
104                        dotted.push(b'.');
105                    }
106                    dotted.extend_from_slice(decimal_u8(octet, &mut buf));
107                }
108                (b"in-addr", dotted.clone(), dotted)
109            }
110            IpAddr::V6(ip) => {
111                let mut segments = Vec::with_capacity(63);
112                for segment in ip.segments() {
113                    for shift in [12u32, 8, 4, 0] {
114                        if !segments.is_empty() {
115                            segments.push(b'.');
116                        }
117                        segments.push(hex_nibble((segment >> shift) as u8));
118                    }
119                }
120                (b"ip6", segments, ip.to_string().into_bytes())
121            }
122        };
123        self.vars[Variable::IpVersion as usize] = v.into();
124        self.vars[Variable::Ip as usize] = i.into();
125        self.vars[Variable::SmtpIp as usize] = c.into();
126    }
127
128    pub fn set_sender(&mut self, value: impl Into<Cow<'x, [u8]>>) {
129        let value = value.into();
130        for (pos, ch) in value.iter().enumerate() {
131            if ch == &b'@' {
132                if pos > 0 {
133                    self.vars[Variable::SenderLocalPart as usize] = match &value {
134                        Cow::Borrowed(value) => (&value[..pos]).into(),
135                        Cow::Owned(value) => value[..pos].to_vec().into(),
136                    };
137                }
138                self.vars[Variable::SenderDomainPart as usize] = match &value {
139                    Cow::Borrowed(value) => (value.get(pos + 1..).unwrap_or_default()).into(),
140                    Cow::Owned(value) => (value.get(pos + 1..).unwrap_or_default()).to_vec().into(),
141                };
142                break;
143            }
144        }
145
146        self.vars[Variable::Sender as usize] = value;
147    }
148
149    pub fn set_helo_domain(&mut self, value: impl Into<Cow<'x, [u8]>>) {
150        self.vars[Variable::HeloDomain as usize] = value.into();
151    }
152
153    pub fn set_host_domain(&mut self, value: impl Into<Cow<'x, [u8]>>) {
154        self.vars[Variable::HostDomain as usize] = value.into();
155    }
156
157    pub fn set_validated_domain(&mut self, value: impl Into<Cow<'x, [u8]>>) {
158        self.vars[Variable::ValidatedDomain as usize] = value.into();
159    }
160
161    pub fn set_domain(&mut self, value: impl Into<Cow<'x, [u8]>>) {
162        self.vars[Variable::Domain as usize] = value.into();
163    }
164
165    pub fn get(
166        &self,
167        name: Variable,
168        num_parts: u32,
169        reverse: bool,
170        escape: bool,
171        fqdn: bool,
172        delimiters: u64,
173    ) -> Cow<'_, [u8]> {
174        let transform = Transform {
175            num_parts,
176            reverse,
177            escape,
178            fqdn,
179            delimiters,
180        };
181        let var: &[u8] = self.vars[name as usize].as_ref();
182        if var.is_empty() && self.current_time_on_demand && matches!(name, Variable::CurrentTime) {
183            let now = current_time();
184            if transform.is_verbatim() {
185                return Cow::Owned(now);
186            }
187            let mut result = Vec::with_capacity(transform.capacity_for(&now));
188            append_transformed(&mut result, &now, transform);
189            return Cow::Owned(result);
190        }
191        if var.is_empty() || transform.is_verbatim() {
192            return Cow::Borrowed(var);
193        }
194
195        let mut result = Vec::with_capacity(transform.capacity_for(var));
196        append_transformed(&mut result, var, transform);
197        Cow::Owned(result)
198    }
199
200    fn append(&self, result: &mut Vec<u8>, name: Variable, transform: Transform) {
201        let var: &[u8] = self.vars[name as usize].as_ref();
202        if var.is_empty() && self.current_time_on_demand && matches!(name, Variable::CurrentTime) {
203            append_variable(result, &current_time(), transform);
204        } else {
205            append_variable(result, var, transform);
206        }
207    }
208}
209
210#[derive(Clone, Copy)]
211struct Transform {
212    num_parts: u32,
213    reverse: bool,
214    escape: bool,
215    fqdn: bool,
216    delimiters: u64,
217}
218
219impl Transform {
220    #[inline(always)]
221    fn is_verbatim(&self) -> bool {
222        self.num_parts == 0
223            && !self.reverse
224            && !self.escape
225            && self.delimiters == DEFAULT_DELIMITERS
226    }
227
228    #[inline(always)]
229    fn capacity_for(&self, var: &[u8]) -> usize {
230        if self.escape {
231            var.len() * 3 + 1
232        } else {
233            var.len() + 1
234        }
235    }
236
237    #[inline(always)]
238    fn is_delimiter(&self, ch: u8) -> bool {
239        let offset = ch.wrapping_sub(b'+');
240        offset <= LAST_DELIMITER && (self.delimiters & (1u64 << offset)) != 0
241    }
242}
243
244fn append_variable(result: &mut Vec<u8>, var: &[u8], transform: Transform) {
245    if var.is_empty() || transform.is_verbatim() {
246        result.extend_from_slice(var);
247    } else {
248        append_transformed(result, var, transform);
249    }
250}
251
252fn append_transformed(result: &mut Vec<u8>, var: &[u8], transform: Transform) {
253    let skipped = if transform.num_parts == 0 {
254        0
255    } else {
256        let total = 1 + var.iter().filter(|ch| transform.is_delimiter(**ch)).count();
257        total - std::cmp::min(total, transform.num_parts as usize)
258    };
259
260    let start = result.len();
261    if !transform.reverse {
262        for (pos, part) in var
263            .split(|ch| transform.is_delimiter(*ch))
264            .skip(skipped)
265            .enumerate()
266        {
267            add_part(result, part, pos, transform.escape);
268        }
269    } else {
270        for (pos, part) in var
271            .rsplit(|ch| transform.is_delimiter(*ch))
272            .skip(skipped)
273            .enumerate()
274        {
275            add_part(result, part, pos, transform.escape);
276        }
277    }
278    if transform.fqdn && !matches!(result.get(start..), Some([.., b'.'])) {
279        result.push(b'.');
280    }
281}
282
283fn current_time() -> Vec<u8> {
284    let mut seconds = SystemTime::now()
285        .duration_since(SystemTime::UNIX_EPOCH)
286        .map(|d| d.as_secs())
287        .unwrap_or(0);
288    let mut buf = [0u8; 20];
289    let mut len = 0;
290    for slot in buf.iter_mut().rev() {
291        *slot = b'0' + (seconds % 10) as u8;
292        len += 1;
293        seconds /= 10;
294        if seconds == 0 {
295            break;
296        }
297    }
298    buf[buf.len() - len..].to_vec()
299}
300
301#[inline(always)]
302fn add_part(result: &mut Vec<u8>, part: &[u8], pos: usize, escape: bool) {
303    if pos > 0 {
304        result.push(b'.');
305    }
306    if !escape {
307        result.extend_from_slice(part);
308    } else {
309        for &ch in part {
310            if ch.is_ascii_alphanumeric() || matches!(ch, b'-' | b'.' | b'_' | b'~') {
311                result.push(ch);
312            } else {
313                result.extend_from_slice(&[b'%', hex_nibble(ch >> 4), hex_nibble(ch)]);
314            }
315        }
316    }
317}
318
319#[cfg(test)]
320mod test {
321    use std::net::IpAddr;
322
323    use crate::spf::{Variables, parse::SPFParser};
324
325    #[test]
326    fn expand_macro() {
327        let mut vars = Variables::new();
328        vars.set_sender("strong-bad@email.example.com".as_bytes());
329        vars.set_ip(&"192.0.2.3".parse::<IpAddr>().unwrap());
330        vars.set_validated_domain("mx.example.org".as_bytes());
331        vars.set_domain("email.example.com".as_bytes());
332        vars.set_helo_domain("....".as_bytes());
333
334        for (macro_string, expansion) in [
335            ("%{s}", "strong-bad@email.example.com"),
336            ("%{o}", "email.example.com"),
337            ("%{d}", "email.example.com"),
338            ("%{d4}", "email.example.com"),
339            ("%{d3}", "email.example.com"),
340            ("%{d2}", "example.com"),
341            ("%{d1}", "com"),
342            ("%{dr}", "com.example.email"),
343            ("%{d2r}", "example.email"),
344            ("%{l}", "strong-bad"),
345            ("%{l-}", "strong.bad"),
346            ("%{lr}", "strong-bad"),
347            ("%{lr-}", "bad.strong"),
348            ("%{l1r-}", "strong"),
349            ("%{p1r}", "mx"),
350            ("%{h3r}", ".."),
351            (
352                "%{ir}.%{v}._spf.%{d2}",
353                "3.2.0.192.in-addr._spf.example.com",
354            ),
355            ("%{lr-}.lp._spf.%{d2}", "bad.strong.lp._spf.example.com"),
356            (
357                "%{lr-}.lp.%{ir}.%{v}._spf.%{d2}",
358                "bad.strong.lp.3.2.0.192.in-addr._spf.example.com",
359            ),
360            (
361                "%{ir}.%{v}.%{l1r-}.lp._spf.%{d2}",
362                "3.2.0.192.in-addr.strong.lp._spf.example.com",
363            ),
364            (
365                "%{d2}.trusted-domains.example.net",
366                "example.com.trusted-domains.example.net",
367            ),
368        ] {
369            let (m, _) = macro_string.as_bytes().iter().macro_string(false).unwrap();
370            assert_eq!(m.eval(&vars, "", false), expansion, "{macro_string:?}");
371        }
372
373        let mut vars = Variables::new();
374        vars.set_sender("strong-bad@email.example.com".as_bytes());
375        vars.set_ip(&"2001:db8::cb01".parse::<IpAddr>().unwrap());
376        vars.set_validated_domain("mx.example.org".as_bytes());
377        vars.set_domain("email.example.com".as_bytes());
378
379        for (macro_string, expansion) in [
380            (
381                "%{ir}.%{v}._spf.%{d2}",
382                concat!(
383                    "1.0.b.c.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.",
384                    "0.0.0.0.0.8.b.d.0.1.0.0.2.ip6._spf.example.com"
385                ),
386            ),
387            ("%{c}", "2001:db8::cb01"),
388            (
389                "%{c} is not one of %{d}'s designated mail servers.",
390                "2001:db8::cb01 is not one of email.example.com's designated mail servers.",
391            ),
392            (
393                "See http://%{d}/why.html?s=%{S}&i=%{C}",
394                concat!(
395                    "See http://email.example.com/why.html?",
396                    "s=strong-bad%40email.example.com&i=2001%3adb8%3a%3acb01"
397                ),
398            ),
399        ] {
400            let (m, _) = macro_string.as_bytes().iter().macro_string(true).unwrap();
401            assert_eq!(m.eval(&vars, "", false), expansion, "{macro_string:?}");
402        }
403    }
404}