Skip to main content

smtp_proto/request/
receiver.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 std::{borrow::Cow, slice::Iter};
8
9use crate::{Error, Request};
10
11pub const MAX_LINE_LENGTH: usize = 4096;
12
13#[derive(Default)]
14pub struct RequestReceiver {
15    buf: Vec<u8>,
16    buf_used: bool,
17}
18
19pub struct DataReceiver {
20    crlf_dot: bool,
21    last_ch: u8,
22    prev_last_ch: u8,
23}
24
25pub struct BdatReceiver {
26    pub is_last: bool,
27    bytes_left: usize,
28}
29
30pub struct DummyDataReceiver {
31    is_bdat: bool,
32    bdat_bytes_left: usize,
33    crlf_dot: bool,
34    last_ch: u8,
35    prev_last_ch: u8,
36}
37
38#[derive(Default)]
39pub struct DummyLineReceiver {}
40
41#[derive(Default)]
42pub struct LineReceiver<T> {
43    pub buf: Vec<u8>,
44    pub state: T,
45}
46
47impl RequestReceiver {
48    pub fn buf(&mut self) -> &mut Vec<u8> {
49        if self.buf_used {
50            self.buf.clear();
51            self.buf_used = false;
52        }
53
54        &mut self.buf
55    }
56
57    pub fn ingest<'this, 'bytes, 'out>(
58        &'this mut self,
59        bytes: &mut Iter<'bytes, u8>,
60    ) -> Result<Request<Cow<'out, str>>, Error>
61    where
62        'this: 'out,
63        'bytes: 'out,
64    {
65        self.buf();
66
67        if self.buf.is_empty() {
68            let buf = bytes.as_slice();
69            match Request::parse(bytes) {
70                Err(Error::NeedsMoreData { bytes_left }) => {
71                    if bytes_left > 0 {
72                        if bytes_left < MAX_LINE_LENGTH {
73                            self.buf = buf[buf.len().saturating_sub(bytes_left)..].to_vec();
74                        } else {
75                            return Err(Error::ResponseTooLong);
76                        }
77                    }
78                }
79                result => return result,
80            }
81        } else {
82            for &ch in bytes {
83                self.buf.push(ch);
84                if ch == b'\n' {
85                    self.buf_used = true;
86                    return Request::parse(&mut self.buf.iter());
87                } else if self.buf.len() == MAX_LINE_LENGTH {
88                    self.buf.clear();
89                    return Err(Error::ResponseTooLong);
90                }
91            }
92        }
93
94        Err(Error::NeedsMoreData { bytes_left: 0 })
95    }
96}
97
98impl DataReceiver {
99    #[allow(clippy::new_without_default)]
100    pub fn new() -> Self {
101        Self {
102            crlf_dot: false,
103            last_ch: 0,
104            prev_last_ch: 0,
105        }
106    }
107
108    pub fn ingest(&mut self, bytes: &mut Iter<'_, u8>, buf: &mut Vec<u8>) -> bool {
109        for &ch in bytes {
110            match ch {
111                b'.' if self.last_ch == b'\n' && self.prev_last_ch == b'\r' => {
112                    self.crlf_dot = true;
113                }
114                b'\n' if self.crlf_dot && self.last_ch == b'\r' => {
115                    buf.truncate(buf.len() - 1);
116                    return true;
117                }
118                b'\r' => {
119                    buf.push(ch);
120                }
121                _ => {
122                    buf.push(ch);
123                    self.crlf_dot = false;
124                }
125            }
126            self.prev_last_ch = self.last_ch;
127            self.last_ch = ch;
128        }
129
130        false
131    }
132}
133
134impl BdatReceiver {
135    pub fn new(chunk_size: usize, is_last: bool) -> Self {
136        Self {
137            bytes_left: chunk_size,
138            is_last,
139        }
140    }
141
142    pub fn ingest(&mut self, bytes: &mut Iter<'_, u8>, buf: &mut Vec<u8>) -> bool {
143        while self.bytes_left > 0 {
144            if let Some(&ch) = bytes.next() {
145                buf.push(ch);
146                self.bytes_left -= 1;
147            } else {
148                return false;
149            }
150        }
151        true
152    }
153}
154
155impl DummyDataReceiver {
156    pub fn new_bdat(chunk_size: usize) -> Self {
157        Self {
158            bdat_bytes_left: chunk_size,
159            is_bdat: true,
160            crlf_dot: false,
161            last_ch: 0,
162            prev_last_ch: 0,
163        }
164    }
165
166    pub fn new_data(data: &DataReceiver) -> Self {
167        Self {
168            is_bdat: false,
169            bdat_bytes_left: 0,
170            crlf_dot: data.crlf_dot,
171            last_ch: data.last_ch,
172            prev_last_ch: data.prev_last_ch,
173        }
174    }
175
176    pub fn ingest(&mut self, bytes: &mut Iter<'_, u8>) -> bool {
177        if !self.is_bdat {
178            for &ch in bytes {
179                match ch {
180                    b'.' if self.last_ch == b'\n' && self.prev_last_ch == b'\r' => {
181                        self.crlf_dot = true;
182                    }
183                    b'\n' if self.crlf_dot && self.last_ch == b'\r' => {
184                        return true;
185                    }
186                    b'\r' => {}
187                    _ => {
188                        self.crlf_dot = false;
189                    }
190                }
191                self.prev_last_ch = self.last_ch;
192                self.last_ch = ch;
193            }
194
195            false
196        } else {
197            while self.bdat_bytes_left > 0 {
198                if bytes.next().is_some() {
199                    self.bdat_bytes_left -= 1;
200                } else {
201                    return false;
202                }
203            }
204
205            true
206        }
207    }
208}
209
210impl<T> LineReceiver<T> {
211    pub fn new(state: T) -> Self {
212        Self {
213            buf: Vec::with_capacity(32),
214            state,
215        }
216    }
217
218    pub fn ingest(&mut self, bytes: &mut Iter<'_, u8>) -> bool {
219        for &ch in bytes {
220            match ch {
221                b'\n' => return true,
222                b'\r' => (),
223                _ => {
224                    if self.buf.len() < MAX_LINE_LENGTH {
225                        self.buf.push(ch);
226                    }
227                }
228            }
229        }
230        false
231    }
232}
233
234impl DummyLineReceiver {
235    pub fn ingest(&mut self, bytes: &mut Iter<'_, u8>) -> bool {
236        for &ch in bytes {
237            if ch == b'\n' {
238                return true;
239            }
240        }
241        false
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::DataReceiver;
248    use crate::{Error, MailFrom, RcptTo, Request, request::receiver::RequestReceiver};
249
250    #[test]
251    fn data_receiver() {
252        'outer: for (data, message) in [
253            (
254                vec!["hi\r\n", "..\r\n", ".a\r\n", "\r\n.\r\n"],
255                "hi\r\n.\r\na\r\n\r\n",
256            ),
257            (
258                vec!["\r\na\rb\nc\r\n.d\r\n..\r\n", "\r\n.\r\n"],
259                "\r\na\rb\nc\r\nd\r\n.\r\n\r\n",
260            ),
261            // Test SMTP smuggling attempts
262            (
263                vec![
264                    "\n.\r\n",
265                    "MAIL FROM:<hello@world.com>\r\n",
266                    "RCPT TO:<test@domain.com\r\n",
267                    "DATA\r\n",
268                    "\r\n.\r\n",
269                ],
270                concat!(
271                    "\n.\r\n",
272                    "MAIL FROM:<hello@world.com>\r\n",
273                    "RCPT TO:<test@domain.com\r\n",
274                    "DATA\r\n",
275                    "\r\n",
276                ),
277            ),
278            (
279                vec![
280                    "\n.\n",
281                    "MAIL FROM:<hello@world.com>\r\n",
282                    "RCPT TO:<test@domain.com\r\n",
283                    "DATA\r\n",
284                    "\r\n.\r\n",
285                ],
286                concat!(
287                    "\n.\n",
288                    "MAIL FROM:<hello@world.com>\r\n",
289                    "RCPT TO:<test@domain.com\r\n",
290                    "DATA\r\n",
291                    "\r\n",
292                ),
293            ),
294            (
295                vec![
296                    "\r.\r\n",
297                    "MAIL FROM:<hello@world.com>\r\n",
298                    "RCPT TO:<test@domain.com\r\n",
299                    "DATA\r\n",
300                    "\r\n.\r\n",
301                ],
302                concat!(
303                    "\r.\r\n",
304                    "MAIL FROM:<hello@world.com>\r\n",
305                    "RCPT TO:<test@domain.com\r\n",
306                    "DATA\r\n",
307                    "\r\n",
308                ),
309            ),
310            (
311                vec![
312                    "\r.\r",
313                    "MAIL FROM:<hello@world.com>\r\n",
314                    "RCPT TO:<test@domain.com\r\n",
315                    "DATA\r\n",
316                    "\r\n.\r\n",
317                ],
318                concat!(
319                    "\r.\r",
320                    "MAIL FROM:<hello@world.com>\r\n",
321                    "RCPT TO:<test@domain.com\r\n",
322                    "DATA\r\n",
323                    "\r\n",
324                ),
325            ),
326        ] {
327            let mut r = DataReceiver::new();
328            let mut buf = Vec::new();
329            for data in &data {
330                if r.ingest(&mut data.as_bytes().iter(), &mut buf) {
331                    assert_eq!(message, String::from_utf8(buf).unwrap());
332                    continue 'outer;
333                }
334            }
335            panic!("Failed for {data:?}");
336        }
337    }
338
339    #[test]
340    fn request_receiver() {
341        for (data, expected_requests) in [
342            (
343                vec![
344                    "data\n",
345                    "start",
346                    "tls\n",
347                    "quit\nnoop",
348                    " hello\nehlo test\nvrfy name\n",
349                    "mail from:<hello",
350                    "@world.com>\nrcpt to:<",
351                    "test@domain.com>\n",
352                ],
353                vec![
354                    Request::Data,
355                    Request::StartTls,
356                    Request::Quit,
357                    Request::Noop {
358                        value: "hello".to_string(),
359                    },
360                    Request::Ehlo {
361                        host: "test".to_string(),
362                    },
363                    Request::Vrfy {
364                        value: "name".to_string(),
365                    },
366                    Request::Mail {
367                        from: MailFrom {
368                            address: "hello@world.com".to_string(),
369                            flags: 0,
370                            size: 0,
371                            trans_id: None,
372                            by: 0,
373                            env_id: None,
374                            solicit: None,
375                            mtrk: None,
376                            auth: None,
377                            hold_for: 0,
378                            hold_until: 0,
379                            mt_priority: 0,
380                        },
381                    },
382                    Request::Rcpt {
383                        to: RcptTo {
384                            address: "test@domain.com".to_string(),
385                            orcpt: None,
386                            rrvs: 0,
387                            flags: 0,
388                        },
389                    },
390                ],
391            ),
392            (
393                vec!["d", "a", "t", "a", "\n", "quit", "\n"],
394                vec![Request::Data, Request::Quit],
395            ),
396        ] {
397            let mut requests = Vec::new();
398            let mut r = RequestReceiver::default();
399            for data in &data {
400                let mut bytes = data.as_bytes().iter();
401                loop {
402                    match r.ingest(&mut bytes) {
403                        Ok(request) => {
404                            requests.push(request.into_owned());
405                            continue;
406                        }
407                        Err(Error::NeedsMoreData { .. }) => {
408                            break;
409                        }
410                        err => panic!("Unexpected error for {data:?}: {err:?}"),
411                    }
412                }
413            }
414            assert_eq!(expected_requests, requests);
415        }
416    }
417}