Skip to main content

mail_auth/common/
headers.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 memchr::{memchr, memchr2};
8
9impl<'x, T> Header<'x, T> {
10    pub fn new(name: &'x [u8], value: &'x [u8], header: T) -> Self {
11        Header {
12            name,
13            value,
14            header,
15        }
16    }
17}
18
19pub trait HeaderStream<'x> {
20    fn next_header(&mut self) -> Option<(&'x [u8], &'x [u8])>;
21    fn body(&mut self) -> &'x [u8];
22}
23
24pub(crate) const MAX_HEADER_LINE_LEN: usize = 76;
25const MEMCHR_MIN_LEN: usize = 16;
26
27pub struct HeaderFolder<'x, W: Writer> {
28    writer: &'x mut W,
29    bytes_left: usize,
30}
31
32impl<'x, W: Writer> HeaderFolder<'x, W> {
33    pub fn new(writer: &'x mut W) -> Self {
34        HeaderFolder {
35            writer,
36            bytes_left: MAX_HEADER_LINE_LEN,
37        }
38    }
39
40    #[inline(always)]
41    fn write_chunk(&mut self, chunk: &[u8]) {
42        if chunk == b"\r\n" {
43            self.writer.write(chunk);
44            self.bytes_left = MAX_HEADER_LINE_LEN;
45        } else if chunk.len() < self.bytes_left {
46            self.writer.write(chunk);
47            self.bytes_left -= chunk.len();
48        } else if chunk.len() >= MAX_HEADER_LINE_LEN {
49            let mut add_new_line = self.bytes_left != MAX_HEADER_LINE_LEN;
50            let mut last_piece_len = MAX_HEADER_LINE_LEN;
51            for chunk in chunk.chunks(MAX_HEADER_LINE_LEN) {
52                if add_new_line {
53                    self.writer.write(b"\r\n\t");
54                }
55                add_new_line = true;
56                self.writer.write(chunk);
57                last_piece_len = chunk.len();
58            }
59            self.bytes_left = MAX_HEADER_LINE_LEN - last_piece_len;
60        } else {
61            self.writer.write(b"\r\n\t");
62            self.writer.write(chunk);
63            self.bytes_left = MAX_HEADER_LINE_LEN - chunk.len();
64        }
65    }
66}
67
68#[inline(always)]
69fn find_semicolon(buf: &[u8]) -> Option<usize> {
70    if buf.len() >= MEMCHR_MIN_LEN {
71        memchr::memchr(b';', buf)
72    } else {
73        buf.iter().position(|ch| *ch == b';')
74    }
75}
76
77impl<'x, W: Writer> Writer for HeaderFolder<'x, W> {
78    fn write(&mut self, buf: &[u8]) {
79        let mut rest = buf;
80        while !rest.is_empty() {
81            let (chunk, tail) = match find_semicolon(rest) {
82                Some(pos) => rest.split_at(pos + 1),
83                None => (rest, Default::default()),
84            };
85            self.write_chunk(chunk);
86            rest = tail;
87        }
88    }
89
90    fn write_chunked(&mut self, buf: &[u8], chunk_len: usize) {
91        if !(3..MAX_HEADER_LINE_LEN).contains(&chunk_len) || find_semicolon(buf).is_some() {
92            for chunk in buf.chunks(chunk_len.max(1)) {
93                self.write(chunk);
94            }
95            return;
96        }
97
98        let mut rest = buf;
99        while rest.len() >= chunk_len {
100            let whole_chunks = self.bytes_left.saturating_sub(1) / chunk_len;
101            if whole_chunks == 0 {
102                self.writer.write(b"\r\n\t");
103                self.bytes_left = MAX_HEADER_LINE_LEN;
104                continue;
105            }
106            let take = (whole_chunks * chunk_len).min(rest.len() - rest.len() % chunk_len);
107            let (head, tail) = rest.split_at(take);
108            self.writer.write(head);
109            self.bytes_left -= take;
110            rest = tail;
111        }
112
113        if !rest.is_empty() {
114            self.write(rest);
115        }
116    }
117}
118
119pub(crate) struct ChainedHeaderIterator<'x, T: Iterator<Item = &'x [u8]>> {
120    parts: T,
121    iter: HeaderIterator<'x>,
122}
123
124pub(crate) struct HeaderIterator<'x> {
125    message: &'x [u8],
126    pos: usize,
127    start_pos: usize,
128}
129
130pub(crate) struct HeaderParser<'x> {
131    message: &'x [u8],
132    pos: usize,
133    start_pos: usize,
134    pub num_received: usize,
135    pub has_message_id: bool,
136    pub has_date: bool,
137}
138
139enum FieldScan {
140    Named { colon: usize, end: usize },
141    Unnamed { end: usize },
142    End { pos: usize },
143}
144
145#[inline(always)]
146fn scan_field(message: &[u8], start_pos: usize, from: usize) -> FieldScan {
147    let mut cur = from;
148
149    while let Some(rest) = message.get(cur..) {
150        let Some(offset) = memchr2(b':', b'\n', rest) else {
151            break;
152        };
153        let Some((head, tail)) = rest.split_at_checked(offset) else {
154            break;
155        };
156        let pos = cur + offset;
157
158        if tail.first() == Some(&b':') {
159            return scan_value(message, pos);
160        } else if head.last() == Some(&b'\r') || pos == start_pos {
161            return FieldScan::End { pos: pos + 1 };
162        }
163
164        match message.get(pos + 1) {
165            Some(b' ' | b'\t') => cur = pos + 1,
166            _ => return FieldScan::Unnamed { end: pos + 1 },
167        }
168    }
169
170    FieldScan::End { pos: message.len() }
171}
172
173#[inline(always)]
174fn scan_value(message: &[u8], colon: usize) -> FieldScan {
175    let mut cur = colon + 1;
176
177    while let Some(rest) = message.get(cur..) {
178        let Some(offset) = memchr(b'\n', rest) else {
179            break;
180        };
181        let pos = cur + offset;
182
183        match message.get(pos + 1) {
184            Some(b' ' | b'\t') => cur = pos + 1,
185            _ => {
186                return FieldScan::Named {
187                    colon,
188                    end: pos + 1,
189                };
190            }
191        }
192    }
193
194    FieldScan::End { pos: message.len() }
195}
196
197#[derive(Debug, Clone, Copy, PartialEq, Eq)]
198pub(crate) enum AuthenticatedHeader<'x> {
199    Ds(&'x [u8]),
200    D2s(&'x [u8]),
201    D2i(&'x [u8]),
202    #[cfg(feature = "arc")]
203    Aar(&'x [u8]),
204    #[cfg(feature = "arc")]
205    Ams(&'x [u8]),
206    #[cfg(feature = "arc")]
207    As(&'x [u8]),
208    From(&'x [u8]),
209    Other(&'x [u8]),
210}
211
212#[derive(Debug, Clone, PartialEq, Eq)]
213pub struct Header<'x, T> {
214    pub name: &'x [u8],
215    pub value: &'x [u8],
216    pub header: T,
217}
218
219impl<'x> HeaderParser<'x> {
220    pub fn new(message: &'x [u8]) -> Self {
221        HeaderParser {
222            message,
223            pos: 0,
224            start_pos: 0,
225            num_received: 0,
226            has_message_id: false,
227            has_date: false,
228        }
229    }
230
231    pub fn body_offset(&mut self) -> Option<usize> {
232        (self.pos < self.message.len()).then_some(self.pos)
233    }
234}
235
236impl<'x> HeaderIterator<'x> {
237    pub fn new(message: &'x [u8]) -> Self {
238        HeaderIterator {
239            message,
240            pos: 0,
241            start_pos: 0,
242        }
243    }
244
245    pub fn seek_start(&mut self) {
246        let rest = self.message.get(self.pos..).unwrap_or_default();
247        self.pos += rest
248            .iter()
249            .position(|ch| !ch.is_ascii_whitespace())
250            .unwrap_or(rest.len());
251    }
252
253    pub fn body_offset(&mut self) -> Option<usize> {
254        (self.pos < self.message.len()).then_some(self.pos)
255    }
256}
257
258impl<'x> HeaderStream<'x> for HeaderIterator<'x> {
259    fn next_header(&mut self) -> Option<(&'x [u8], &'x [u8])> {
260        self.next()
261    }
262
263    fn body(&mut self) -> &'x [u8] {
264        self.body_offset()
265            .and_then(|offset| self.message.get(offset..))
266            .unwrap_or_default()
267    }
268}
269
270impl<'x> Iterator for HeaderIterator<'x> {
271    type Item = (&'x [u8], &'x [u8]);
272
273    fn next(&mut self) -> Option<Self::Item> {
274        match scan_field(self.message, self.start_pos, self.pos) {
275            FieldScan::Named { colon, end } => {
276                let header_name = self.message.get(self.start_pos..colon).unwrap_or_default();
277                let header_value = self.message.get(colon + 1..end).unwrap_or_default();
278
279                self.start_pos = end;
280                self.pos = end;
281
282                Some((header_name, header_value))
283            }
284            FieldScan::Unnamed { end } => {
285                let header_name = self.message.get(self.start_pos..end).unwrap_or_default();
286
287                self.start_pos = end;
288                self.pos = end;
289
290                Some((header_name, b""))
291            }
292            FieldScan::End { pos } => {
293                self.pos = pos;
294
295                None
296            }
297        }
298    }
299}
300
301impl<'x, T: Iterator<Item = &'x [u8]>> ChainedHeaderIterator<'x, T> {
302    pub fn new(mut parts: T) -> Self {
303        ChainedHeaderIterator {
304            iter: HeaderIterator::new(parts.next().unwrap_or_default()),
305            parts,
306        }
307    }
308}
309
310impl<'x, T: Iterator<Item = &'x [u8]>> HeaderStream<'x> for ChainedHeaderIterator<'x, T> {
311    fn next_header(&mut self) -> Option<(&'x [u8], &'x [u8])> {
312        if let Some(header) = self.iter.next_header() {
313            Some(header)
314        } else {
315            self.iter = HeaderIterator::new(self.parts.next()?);
316            self.iter.next_header()
317        }
318    }
319
320    fn body(&mut self) -> &'x [u8] {
321        self.iter.body()
322    }
323}
324
325impl<'x> Iterator for HeaderParser<'x> {
326    type Item = (AuthenticatedHeader<'x>, &'x [u8]);
327
328    fn next(&mut self) -> Option<Self::Item> {
329        let scan_start = self.pos;
330
331        match scan_field(self.message, self.start_pos, scan_start) {
332            FieldScan::Named { colon, end } => {
333                let header_name = self.message.get(self.start_pos..colon).unwrap_or_default();
334                let header_value = self.message.get(colon + 1..end).unwrap_or_default();
335                let token = self.message.get(scan_start..colon).unwrap_or_default();
336
337                self.start_pos = end;
338                self.pos = end;
339
340                let header_name = match classify_name(token) {
341                    HeaderClass::Received => {
342                        self.num_received += 1;
343                        AuthenticatedHeader::Other(header_name)
344                    }
345                    HeaderClass::MessageId => {
346                        self.has_message_id = true;
347                        AuthenticatedHeader::Other(header_name)
348                    }
349                    HeaderClass::Date => {
350                        self.has_date = true;
351                        AuthenticatedHeader::Other(header_name)
352                    }
353                    HeaderClass::From => AuthenticatedHeader::From(header_name),
354                    HeaderClass::Ds => AuthenticatedHeader::Ds(header_name),
355                    HeaderClass::D2s => AuthenticatedHeader::D2s(header_name),
356                    HeaderClass::D2i => AuthenticatedHeader::D2i(header_name),
357                    #[cfg(feature = "arc")]
358                    HeaderClass::Aar => AuthenticatedHeader::Aar(header_name),
359                    #[cfg(feature = "arc")]
360                    HeaderClass::Ams => AuthenticatedHeader::Ams(header_name),
361                    #[cfg(feature = "arc")]
362                    HeaderClass::As => AuthenticatedHeader::As(header_name),
363                    #[cfg(not(feature = "arc"))]
364                    HeaderClass::Aar | HeaderClass::Ams | HeaderClass::As => {
365                        AuthenticatedHeader::Other(header_name)
366                    }
367                    HeaderClass::Other => AuthenticatedHeader::Other(header_name),
368                };
369
370                Some((header_name, header_value))
371            }
372            FieldScan::Unnamed { end } => {
373                let header_name = self.message.get(self.start_pos..end).unwrap_or_default();
374
375                self.start_pos = end;
376                self.pos = end;
377
378                Some((AuthenticatedHeader::Other(header_name), b""))
379            }
380            FieldScan::End { pos } => {
381                self.pos = pos;
382
383                None
384            }
385        }
386    }
387}
388
389enum HeaderClass {
390    Received,
391    From,
392    Date,
393    MessageId,
394    Ds,
395    D2s,
396    D2i,
397    Aar,
398    Ams,
399    As,
400    Other,
401}
402
403#[inline(always)]
404fn is_field_token(ch: u8) -> bool {
405    ch.is_ascii_alphanumeric() || ch == b'-'
406}
407
408#[inline(always)]
409fn classify_name(name: &[u8]) -> HeaderClass {
410    if name.iter().fold(true, |acc, &ch| acc & is_field_token(ch)) {
411        match name.len() {
412            4 if name.eq_ignore_ascii_case(b"from") => HeaderClass::From,
413            4 if name.eq_ignore_ascii_case(b"date") => HeaderClass::Date,
414            8 if name.eq_ignore_ascii_case(b"received") => HeaderClass::Received,
415            10 if name.eq_ignore_ascii_case(b"message-id") => HeaderClass::MessageId,
416            14 if name.eq_ignore_ascii_case(b"dkim-signature") => HeaderClass::Ds,
417            15 if name.eq_ignore_ascii_case(b"dkim2-signature") => HeaderClass::D2s,
418            16 if name.eq_ignore_ascii_case(b"message-instance") => HeaderClass::D2i,
419            21 if name.eq_ignore_ascii_case(b"arc-message-signature") => HeaderClass::Ams,
420            26 if name.eq_ignore_ascii_case(b"arc-authentication-results") => HeaderClass::Aar,
421            _ if name
422                .get(..8)
423                .is_some_and(|prefix| prefix.eq_ignore_ascii_case(b"arc-seal")) =>
424            {
425                HeaderClass::As
426            }
427            _ => HeaderClass::Other,
428        }
429    } else {
430        classify_folded_name(name)
431    }
432}
433
434#[inline(never)]
435fn classify_folded_name(name: &[u8]) -> HeaderClass {
436    let mut token_start = usize::MAX;
437    let mut token_end = usize::MAX;
438
439    let mut hash: u64 = 0;
440    let mut hash_shift = 0;
441
442    for (pos, &ch) in name.iter().enumerate() {
443        let token = match ch {
444            b' ' | b'\t' | b'\r' | b'\n' => continue,
445            b'A'..=b'Z' => ch - b'A' + b'a',
446            b'a'..=b'z' | b'-' | b'0'..=b'9' => ch,
447            _ => {
448                hash = u64::MAX;
449                continue;
450            }
451        };
452
453        if hash_shift < 64 {
454            hash |= (token as u64) << hash_shift;
455            hash_shift += 8;
456
457            if token_start == usize::MAX {
458                token_start = pos;
459            }
460        }
461        token_end = pos;
462    }
463
464    let tail = name
465        .get(token_start.wrapping_add(8)..token_end.wrapping_add(1))
466        .unwrap_or_default();
467
468    match hash {
469        RECEIVED if token_start.wrapping_add(7) == token_end => HeaderClass::Received,
470        FROM => HeaderClass::From,
471        AS => HeaderClass::As,
472        AAR if tail.eq_ignore_ascii_case(b"entication-Results") => HeaderClass::Aar,
473        AMS if tail.eq_ignore_ascii_case(b"age-Signature") => HeaderClass::Ams,
474        DKIM if tail.eq_ignore_ascii_case(b"nature") => HeaderClass::Ds,
475        DKIM2 if tail.eq_ignore_ascii_case(b"gnature") => HeaderClass::D2s,
476        MSGID if tail.eq_ignore_ascii_case(b"id") => HeaderClass::MessageId,
477        MSGID if tail.eq_ignore_ascii_case(b"instance") => HeaderClass::D2i,
478        DATE => HeaderClass::Date,
479        _ => HeaderClass::Other,
480    }
481}
482
483pub(crate) const HEADER_CAPACITY: usize = 512;
484
485pub trait HeaderWriter: Sized {
486    fn write_header(&self, writer: &mut impl Writer);
487    fn to_header(&self) -> String {
488        let mut buf = Vec::with_capacity(HEADER_CAPACITY);
489        self.write_header(&mut buf);
490        String::from_utf8(buf)
491            .unwrap_or_else(|err| String::from_utf8_lossy(err.as_bytes()).into_owned())
492    }
493}
494
495pub trait Writable {
496    fn write(self, writer: &mut impl Writer);
497}
498
499impl Writable for &[u8] {
500    fn write(self, writer: &mut impl Writer) {
501        writer.write(self);
502    }
503}
504
505pub trait Writer {
506    fn write(&mut self, buf: &[u8]);
507
508    fn write_len(&mut self, buf: &[u8], len: &mut usize) {
509        self.write(buf);
510        *len += buf.len();
511    }
512
513    /// Writes `buf` as if it had been split into `chunk_len` sized pieces, each
514    /// passed to [`Writer::write`] in turn. Writers whose output only depends on
515    /// the concatenation of what they receive take the default implementation.
516    fn write_chunked(&mut self, buf: &[u8], _chunk_len: usize) {
517        self.write(buf);
518    }
519}
520
521impl Writer for Vec<u8> {
522    fn write(&mut self, buf: &[u8]) {
523        self.extend(buf);
524    }
525}
526
527impl Writer for &mut Vec<u8> {
528    fn write(&mut self, buf: &[u8]) {
529        self.extend(buf);
530    }
531}
532
533const MAX_U64_DIGITS: usize = 20;
534
535pub(crate) struct IntegerBuffer([u8; MAX_U64_DIGITS]);
536
537impl IntegerBuffer {
538    pub(crate) const fn new() -> Self {
539        IntegerBuffer([0; MAX_U64_DIGITS])
540    }
541
542    pub(crate) fn digits(&mut self, value: u64) -> &[u8] {
543        let mut value = value;
544        let mut pos = MAX_U64_DIGITS;
545        loop {
546            pos -= 1;
547            if let Some(digit) = self.0.get_mut(pos) {
548                *digit = b'0' + (value % 10) as u8;
549            }
550            value /= 10;
551            if value == 0 || pos == 0 {
552                break;
553            }
554        }
555        self.0.get(pos..).unwrap_or_default()
556    }
557
558    pub(crate) fn text(&mut self, value: u64) -> &str {
559        std::str::from_utf8(self.digits(value)).unwrap_or_default()
560    }
561}
562
563pub(crate) fn write_integer(writer: &mut impl Writer, value: u64) {
564    let mut buffer = IntegerBuffer::new();
565    writer.write(buffer.digits(value));
566}
567
568pub(crate) fn write_wrapped(
569    writer: &mut impl Writer,
570    value: &[u8],
571    bytes_written: &mut usize,
572    new_line: &[u8],
573) {
574    let mut rest = value;
575    while !rest.is_empty() {
576        let take = MAX_HEADER_LINE_LEN
577            .saturating_sub(*bytes_written)
578            .max(1)
579            .min(rest.len());
580        let (head, tail) = rest.split_at(take);
581        writer.write_len(head, bytes_written);
582        if *bytes_written >= MAX_HEADER_LINE_LEN {
583            writer.write(new_line);
584            *bytes_written = 1;
585        }
586        rest = tail;
587    }
588}
589
590const BASE64_ALPHABET: &[u8; 64] =
591    b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
592pub(crate) const BASE64_GROUP_LEN: usize = 4;
593const BASE64_INPUT_LEN: usize = 192;
594const BASE64_OUTPUT_LEN: usize = BASE64_INPUT_LEN / 3 * BASE64_GROUP_LEN;
595
596#[inline(always)]
597fn base64_group(word: u32, len: usize) -> [u8; BASE64_GROUP_LEN] {
598    let pad = |shift: u32, keep: usize| {
599        if len > keep {
600            BASE64_ALPHABET[(word >> shift) as usize & 0x3f]
601        } else {
602            b'='
603        }
604    };
605    [
606        BASE64_ALPHABET[(word >> 18) as usize & 0x3f],
607        BASE64_ALPHABET[(word >> 12) as usize & 0x3f],
608        pad(6, 1),
609        pad(0, 2),
610    ]
611}
612
613pub(crate) fn base64_encode_slice(bytes: &[u8], out: &mut [u8]) -> usize {
614    let (groups, tail) = bytes.as_chunks::<3>();
615    let mut written = 0;
616
617    let (slots, _) = out.as_chunks_mut::<BASE64_GROUP_LEN>();
618    for (group, slot) in groups.iter().zip(slots.iter_mut()) {
619        let [b0, b1, b2] = *group;
620        let word = ((b0 as u32) << 16) | ((b1 as u32) << 8) | b2 as u32;
621        slot.copy_from_slice(&base64_group(word, 3));
622        written += BASE64_GROUP_LEN;
623    }
624
625    if !tail.is_empty() {
626        let word = tail.iter().enumerate().fold(0u32, |word, (pos, byte)| {
627            word | (*byte as u32) << (16 - pos * 8)
628        });
629        if let Some(slot) = out.get_mut(written..written + BASE64_GROUP_LEN) {
630            slot.copy_from_slice(&base64_group(word, tail.len()));
631            written += BASE64_GROUP_LEN;
632        }
633    }
634
635    written
636}
637
638pub(crate) fn write_base64(writer: &mut impl Writer, bytes: &[u8]) {
639    let mut buffer = [0u8; BASE64_OUTPUT_LEN];
640    for window in bytes.chunks(BASE64_INPUT_LEN) {
641        let written = base64_encode_slice(window, &mut buffer);
642        writer.write_chunked(buffer.get(..written).unwrap_or_default(), BASE64_GROUP_LEN);
643    }
644}
645
646pub(crate) fn write_wrapped_base64(
647    writer: &mut impl Writer,
648    bytes: &[u8],
649    bytes_written: &mut usize,
650    new_line: &[u8],
651) {
652    let mut buffer = [0u8; BASE64_OUTPUT_LEN];
653    for window in bytes.chunks(BASE64_INPUT_LEN) {
654        let written = base64_encode_slice(window, &mut buffer);
655        write_wrapped(
656            writer,
657            buffer.get(..written).unwrap_or_default(),
658            bytes_written,
659            new_line,
660        );
661    }
662}
663
664const FROM: u64 =
665    (b'f' as u64) | ((b'r' as u64) << 8) | ((b'o' as u64) << 16) | ((b'm' as u64) << 24);
666const DKIM: u64 = (b'd' as u64)
667    | ((b'k' as u64) << 8)
668    | ((b'i' as u64) << 16)
669    | ((b'm' as u64) << 24)
670    | ((b'-' as u64) << 32)
671    | ((b's' as u64) << 40)
672    | ((b'i' as u64) << 48)
673    | ((b'g' as u64) << 56);
674const DKIM2: u64 = (b'd' as u64)
675    | ((b'k' as u64) << 8)
676    | ((b'i' as u64) << 16)
677    | ((b'm' as u64) << 24)
678    | ((b'2' as u64) << 32)
679    | ((b'-' as u64) << 40)
680    | ((b's' as u64) << 48)
681    | ((b'i' as u64) << 56);
682const AAR: u64 = (b'a' as u64)
683    | ((b'r' as u64) << 8)
684    | ((b'c' as u64) << 16)
685    | ((b'-' as u64) << 24)
686    | ((b'a' as u64) << 32)
687    | ((b'u' as u64) << 40)
688    | ((b't' as u64) << 48)
689    | ((b'h' as u64) << 56);
690const AMS: u64 = (b'a' as u64)
691    | ((b'r' as u64) << 8)
692    | ((b'c' as u64) << 16)
693    | ((b'-' as u64) << 24)
694    | ((b'm' as u64) << 32)
695    | ((b'e' as u64) << 40)
696    | ((b's' as u64) << 48)
697    | ((b's' as u64) << 56);
698const AS: u64 = (b'a' as u64)
699    | ((b'r' as u64) << 8)
700    | ((b'c' as u64) << 16)
701    | ((b'-' as u64) << 24)
702    | ((b's' as u64) << 32)
703    | ((b'e' as u64) << 40)
704    | ((b'a' as u64) << 48)
705    | ((b'l' as u64) << 56);
706const RECEIVED: u64 = (b'r' as u64)
707    | ((b'e' as u64) << 8)
708    | ((b'c' as u64) << 16)
709    | ((b'e' as u64) << 24)
710    | ((b'i' as u64) << 32)
711    | ((b'v' as u64) << 40)
712    | ((b'e' as u64) << 48)
713    | ((b'd' as u64) << 56);
714const DATE: u64 =
715    (b'd' as u64) | ((b'a' as u64) << 8) | ((b't' as u64) << 16) | ((b'e' as u64) << 24);
716const MSGID: u64 = (b'm' as u64)
717    | ((b'e' as u64) << 8)
718    | ((b's' as u64) << 16)
719    | ((b's' as u64) << 24)
720    | ((b'a' as u64) << 32)
721    | ((b'g' as u64) << 40)
722    | ((b'e' as u64) << 48)
723    | ((b'-' as u64) << 56);
724
725#[cfg(test)]
726mod test {
727    use super::{ChainedHeaderIterator, HeaderIterator, HeaderStream};
728    use super::{HeaderFolder, MAX_HEADER_LINE_LEN};
729    use crate::common::headers::{AuthenticatedHeader, HeaderParser, Writer};
730
731    #[test]
732    fn header_iterator() {
733        for (message, headers) in [
734            (
735                "From: a\nTo: b\nEmpty:\nMulti: 1\n 2\nSubject: c\n\nNot-header: ignore\n",
736                vec![
737                    ("From", " a\n"),
738                    ("To", " b\n"),
739                    ("Empty", "\n"),
740                    ("Multi", " 1\n 2\n"),
741                    ("Subject", " c\n"),
742                ],
743            ),
744            (
745                ": a\nTo: b\n \n \nc\n:\nFrom : d\nSubject: e\n\nNot-header: ignore\n",
746                vec![
747                    ("", " a\n"),
748                    ("To", " b\n \n \n"),
749                    ("c\n", ""),
750                    ("", "\n"),
751                    ("From ", " d\n"),
752                    ("Subject", " e\n"),
753                ],
754            ),
755            (
756                concat!(
757                    "A: X\r\n",
758                    "B : Y\t\r\n",
759                    "\tZ  \r\n",
760                    "\r\n",
761                    " C \r\n",
762                    "D \t E\r\n"
763                ),
764                vec![("A", " X\r\n"), ("B ", " Y\t\r\n\tZ  \r\n")],
765            ),
766        ] {
767            assert_eq!(
768                HeaderIterator::new(message.as_bytes())
769                    .map(|(h, v)| {
770                        (
771                            std::str::from_utf8(h).unwrap(),
772                            std::str::from_utf8(v).unwrap(),
773                        )
774                    })
775                    .collect::<Vec<_>>(),
776                headers
777            );
778
779            assert_eq!(
780                HeaderParser::new(message.as_bytes())
781                    .map(|(h, v)| {
782                        (
783                            std::str::from_utf8(match h {
784                                #[cfg(feature = "arc")]
785                                AuthenticatedHeader::Aar(v)
786                                | AuthenticatedHeader::Ams(v)
787                                | AuthenticatedHeader::As(v) => v,
788                                AuthenticatedHeader::Ds(v)
789                                | AuthenticatedHeader::D2s(v)
790                                | AuthenticatedHeader::D2i(v)
791                                | AuthenticatedHeader::From(v)
792                                | AuthenticatedHeader::Other(v) => v,
793                            })
794                            .unwrap(),
795                            std::str::from_utf8(v).unwrap(),
796                        )
797                    })
798                    .collect::<Vec<_>>(),
799                headers
800            );
801        }
802    }
803
804    #[cfg(feature = "arc")]
805    #[test]
806    fn header_parser() {
807        let message = concat!(
808            "ARC-Message-Signature: i=1; a=rsa-sha256;\n",
809            "ARC-Authentication-Results: i=1;\n",
810            "ARC-Seal: i=1; a=rsa-sha256;\n",
811            "DKIM-Signature: v=1; a=rsa-sha256; c=relaxed/simple;\n",
812            "From: jdoe@domain\n",
813            "F r o m : jane@domain.com\n",
814            "ARC-Authentication: i=1;\n",
815            "Received: r1\n",
816            "Received: r2\n",
817            "Received: r3\n",
818            "Received-From: test\n",
819            "Date: date\n",
820            "Message-Id: myid\n",
821            "\nhey",
822        );
823        let mut parser = HeaderParser::new(message.as_bytes());
824        assert_eq!(
825            (&mut parser).map(|(h, _)| { h }).collect::<Vec<_>>(),
826            vec![
827                AuthenticatedHeader::Ams(b"ARC-Message-Signature"),
828                AuthenticatedHeader::Aar(b"ARC-Authentication-Results"),
829                AuthenticatedHeader::As(b"ARC-Seal"),
830                AuthenticatedHeader::Ds(b"DKIM-Signature"),
831                AuthenticatedHeader::From(b"From"),
832                AuthenticatedHeader::From(b"F r o m "),
833                AuthenticatedHeader::Other(b"ARC-Authentication"),
834                AuthenticatedHeader::Other(b"Received"),
835                AuthenticatedHeader::Other(b"Received"),
836                AuthenticatedHeader::Other(b"Received"),
837                AuthenticatedHeader::Other(b"Received-From"),
838                AuthenticatedHeader::Other(b"Date"),
839                AuthenticatedHeader::Other(b"Message-Id"),
840            ]
841        );
842        assert!(parser.has_date);
843        assert!(parser.has_message_id);
844        assert_eq!(parser.num_received, 3);
845    }
846
847    #[test]
848    fn chained_header_iterator() {
849        let parts = [
850            &b"From: a\nTo: b\nEmpty:\nMulti: 1\n 2\n"[..],
851            &b"Subject: c\nReceived: d\n\nhey"[..],
852        ];
853        let mut headers = vec![
854            ("From", " a\n"),
855            ("To", " b\n"),
856            ("Empty", "\n"),
857            ("Multi", " 1\n 2\n"),
858            ("Subject", " c\n"),
859            ("Received", " d\n"),
860        ]
861        .into_iter();
862        let mut it = ChainedHeaderIterator::new(parts.iter().copied());
863
864        while let Some((k, v)) = it.next_header() {
865            assert_eq!(
866                (
867                    std::str::from_utf8(k).unwrap(),
868                    std::str::from_utf8(v).unwrap()
869                ),
870                headers.next().unwrap()
871            );
872        }
873        assert_eq!(it.body(), b"hey");
874    }
875
876    fn fold(header: &[u8]) -> Vec<u8> {
877        let mut buf = Vec::with_capacity(header.len() + 16);
878        let mut folder = HeaderFolder::new(&mut buf);
879        folder.write(header);
880        buf
881    }
882
883    fn unfold(folded: &[u8]) -> Vec<u8> {
884        let mut out = Vec::with_capacity(folded.len());
885        let mut i = 0;
886        while i < folded.len() {
887            if folded[i..].starts_with(b"\r\n\t") {
888                i += 3;
889            } else {
890                out.push(folded[i]);
891                i += 1;
892            }
893        }
894        out
895    }
896
897    fn assert_folded(original: &[u8]) -> Vec<u8> {
898        let folded = fold(original);
899
900        assert_eq!(
901            unfold(&folded),
902            original,
903            "folding must only insert CRLF+TAB fold points, never alter bytes: {:?}",
904            String::from_utf8_lossy(original)
905        );
906
907        for (n, line) in folded.split(|&c| c == b'\n').enumerate() {
908            let line = line.strip_suffix(b"\r").unwrap_or(line);
909            let content = line.strip_prefix(b"\t").unwrap_or(line);
910            assert!(
911                content.len() <= MAX_HEADER_LINE_LEN,
912                "line {n} of {content_len} bytes exceeds {MAX_HEADER_LINE_LEN}: {:?}",
913                String::from_utf8_lossy(content),
914                content_len = content.len(),
915            );
916        }
917
918        for (i, &ch) in folded.iter().enumerate() {
919            if ch == b'\n' {
920                assert!(
921                    i >= 1 && folded[i - 1] == b'\r' && folded.get(i + 1) == Some(&b'\t'),
922                    "every LF must be part of a CRLF+TAB fold at offset {i}: {:?}",
923                    String::from_utf8_lossy(&folded)
924                );
925            }
926        }
927
928        assert!(
929            !folded.starts_with(b"\r\n\t"),
930            "output must never begin with a fold"
931        );
932
933        folded
934    }
935
936    fn extract_header<'a>(eml: &'a str, name: &str) -> &'a [u8] {
937        eml.lines()
938            .find(|l| {
939                l.len() > name.len()
940                    && l.as_bytes()[..name.len()].eq_ignore_ascii_case(name.as_bytes())
941                    && l.as_bytes()[name.len()] == b':'
942            })
943            .unwrap_or_else(|| panic!("header {name} not found"))
944            .as_bytes()
945    }
946
947    #[test]
948    fn header_folder_passthrough() {
949        for input in [
950            &b""[..],
951            &b";"[..],
952            &b"a;;b;;;c"[..],
953            &b"Subject: hello world"[..],
954            &b"Dkim2-Signature: i=1; m=1; d=test.dkim2.eu"[..],
955        ] {
956            assert_eq!(
957                fold(input),
958                input,
959                "should pass through unchanged: {:?}",
960                String::from_utf8_lossy(input)
961            );
962        }
963
964        let just_under = vec![b'a'; MAX_HEADER_LINE_LEN - 1];
965        assert_eq!(fold(&just_under), just_under);
966    }
967
968    #[test]
969    fn header_folder_boundaries() {
970        let exactly_max = vec![b'a'; MAX_HEADER_LINE_LEN];
971        let folded = assert_folded(&exactly_max);
972        assert_eq!(folded, exactly_max, "76 bytes fit on one line, no fold");
973
974        let over_max = vec![b'a'; MAX_HEADER_LINE_LEN + 1];
975        let folded = assert_folded(&over_max);
976        let mut expected = vec![b'a'; MAX_HEADER_LINE_LEN];
977        expected.extend_from_slice(b"\r\n\t");
978        expected.push(b'a');
979        assert_eq!(folded, expected, "77 bytes wrap into 76 + fold + 1");
980
981        let two_pieces = vec![b'a'; MAX_HEADER_LINE_LEN * 2];
982        let folded = assert_folded(&two_pieces);
983        assert_eq!(
984            folded.iter().filter(|&&c| c == b'\n').count(),
985            1,
986            "an exact multiple of the limit yields exactly one fold"
987        );
988    }
989
990    #[test]
991    fn header_folder_large_chunk_followed_by_tags() {
992        let big = vec![b'A'; MAX_HEADER_LINE_LEN * 2 - 3];
993        let mut input = b"v=".to_vec();
994        input.extend_from_slice(&big);
995        input.extend_from_slice(b";a=1;b=2;c=3;d=4;e=5");
996        assert_folded(&input);
997
998        let big = vec![b'A'; MAX_HEADER_LINE_LEN + 20];
999        let mut input = b"s=".to_vec();
1000        input.extend_from_slice(&big);
1001        input.extend_from_slice(b";f=feedback");
1002        assert_folded(&input);
1003    }
1004
1005    #[test]
1006    fn header_folder_consecutive_large_chunks() {
1007        let mf = vec![b'A'; 100];
1008        let rt = vec![b'B'; 90];
1009        let mut input = b"Dkim2-Signature:mf=".to_vec();
1010        input.extend_from_slice(&mf);
1011        input.extend_from_slice(b";rt=");
1012        input.extend_from_slice(&rt);
1013        input.extend_from_slice(b";f=feedback");
1014        assert_folded(&input);
1015    }
1016
1017    #[test]
1018    fn header_folder_many_small_tags() {
1019        let mut input = b"Dkim2-Signature:".to_vec();
1020        for i in 0..40 {
1021            input.extend_from_slice(format!(" tag{i}=value{i};").as_bytes());
1022        }
1023        assert_folded(&input);
1024    }
1025
1026    #[test]
1027    fn header_folder_large_leading_chunk_over_partial_line() {
1028        let mut input = b"Message-Instance: m=1; h=sha256:".to_vec();
1029        input.extend_from_slice(&[b'Z'; 120]);
1030        assert_folded(&input);
1031    }
1032
1033    #[test]
1034    fn header_folder_real_dkim2_headers() {
1035        const FILES: [&str; 2] = [
1036            include_str!("../../resources/dkim2/expected/d2_duplicate_rt_tag.eml"),
1037            include_str!("../../resources/dkim2/expected/pkix_rsa8192.eml"),
1038        ];
1039
1040        for eml in FILES {
1041            for name in ["Message-Instance", "Dkim2-Signature"] {
1042                let header = extract_header(eml, name);
1043                assert!(
1044                    header.len() > MAX_HEADER_LINE_LEN,
1045                    "{name} should exceed the fold limit to exercise folding"
1046                );
1047                let folded = assert_folded(header);
1048                assert!(
1049                    folded.windows(3).any(|w| w == b"\r\n\t"),
1050                    "long real header {name} should have been folded"
1051                );
1052            }
1053        }
1054    }
1055}