1use 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, ¤t_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}