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