Skip to main content

mail_builder/
writer.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::io::{self, Write};
8
9/// Byte sink used by every serializer in this crate.
10///
11/// `Vec<u8>` is the primary implementation; [`IoWriter`] adapts any
12/// [`std::io::Write`] with an internal buffer.
13pub trait Writer {
14    fn write(&mut self, bytes: &[u8]);
15
16    #[inline]
17    fn write_byte(&mut self, byte: u8) {
18        self.write(&[byte]);
19    }
20
21    #[inline]
22    fn reserve(&mut self, _additional: usize) {}
23
24    /// Hands the sink a scratch region of `len` bytes, keeps the first
25    /// `fill(region)` bytes of it and discards the rest.
26    fn write_with(&mut self, len: usize, fill: impl FnOnce(&mut [u8]) -> usize) {
27        let mut buffer = vec![0; len];
28        let written = fill(&mut buffer).min(len);
29        self.write(buffer.get(..written).unwrap_or_default());
30    }
31}
32
33impl Writer for Vec<u8> {
34    #[inline(always)]
35    fn write(&mut self, bytes: &[u8]) {
36        self.extend_from_slice(bytes);
37    }
38
39    #[inline(always)]
40    fn write_byte(&mut self, byte: u8) {
41        self.push(byte);
42    }
43
44    #[inline]
45    fn reserve(&mut self, additional: usize) {
46        Vec::reserve(self, additional);
47    }
48
49    #[inline]
50    fn write_with(&mut self, len: usize, fill: impl FnOnce(&mut [u8]) -> usize) {
51        let start = self.len();
52        self.resize(start + len, 0);
53        let written = fill(self.get_mut(start..).unwrap_or_default()).min(len);
54        self.truncate(start + written);
55    }
56}
57
58impl<W: Writer + ?Sized> Writer for &mut W {
59    #[inline(always)]
60    fn write(&mut self, bytes: &[u8]) {
61        (**self).write(bytes);
62    }
63
64    #[inline(always)]
65    fn write_byte(&mut self, byte: u8) {
66        (**self).write_byte(byte);
67    }
68
69    #[inline]
70    fn reserve(&mut self, additional: usize) {
71        (**self).reserve(additional);
72    }
73
74    #[inline]
75    fn write_with(&mut self, len: usize, fill: impl FnOnce(&mut [u8]) -> usize) {
76        (**self).write_with(len, fill);
77    }
78}
79
80pub(crate) const IO_BUFFER_MAX: usize = 64 * 1024;
81pub(crate) const IO_BUFFER_MIN: usize = 4 * 1024;
82
83/// Buffered adapter that lets any [`std::io::Write`] act as a [`Writer`].
84///
85/// Errors are sticky: the first failure stops all further output and is
86/// returned by [`IoWriter::finish`]. Buffered bytes are only written by
87/// [`IoWriter::finish`] or [`IoWriter::into_result`]; dropping the adapter
88/// discards them.
89pub struct IoWriter<W: Write> {
90    inner: W,
91    buffer: Vec<u8>,
92    error: Option<io::Error>,
93}
94
95impl<W: Write> IoWriter<W> {
96    pub fn new(inner: W) -> Self {
97        Self::with_capacity(IO_BUFFER_MAX, inner)
98    }
99
100    pub fn with_capacity(capacity: usize, inner: W) -> Self {
101        IoWriter {
102            inner,
103            buffer: Vec::with_capacity(capacity.max(1)),
104            error: None,
105        }
106    }
107
108    fn write_inner(&mut self, bytes: &[u8]) {
109        if self.error.is_none()
110            && let Err(err) = self.inner.write_all(bytes)
111        {
112            self.error = Some(err);
113        }
114    }
115
116    fn flush_buffer(&mut self) {
117        if !self.buffer.is_empty() {
118            let buffer = std::mem::take(&mut self.buffer);
119            self.write_inner(&buffer);
120            self.buffer = buffer;
121            self.buffer.clear();
122        }
123    }
124
125    /// Flushes the buffer and returns the wrapped writer, or the first
126    /// error that occurred.
127    pub fn finish(mut self) -> io::Result<W> {
128        self.flush_buffer();
129        match self.error.take() {
130            Some(err) => Err(err),
131            None => Ok(self.inner),
132        }
133    }
134
135    pub fn into_result(self) -> io::Result<()> {
136        self.finish().map(|_| ())
137    }
138}
139
140impl<W: Write> Writer for IoWriter<W> {
141    #[inline]
142    fn write(&mut self, bytes: &[u8]) {
143        if bytes.len() > self.buffer.capacity() - self.buffer.len() {
144            self.flush_buffer();
145            if bytes.len() >= self.buffer.capacity() {
146                self.write_inner(bytes);
147                return;
148            }
149        }
150        self.buffer.extend_from_slice(bytes);
151    }
152
153    #[inline]
154    fn write_byte(&mut self, byte: u8) {
155        if self.buffer.len() == self.buffer.capacity() {
156            self.flush_buffer();
157        }
158        self.buffer.push(byte);
159    }
160
161    #[inline]
162    fn write_with(&mut self, len: usize, fill: impl FnOnce(&mut [u8]) -> usize) {
163        if len > self.buffer.capacity() - self.buffer.len() {
164            self.flush_buffer();
165        }
166        let start = self.buffer.len();
167        self.buffer.resize(start + len, 0);
168        let written = fill(self.buffer.get_mut(start..).unwrap_or_default()).min(len);
169        self.buffer.truncate(start + written);
170    }
171}
172
173#[cfg(test)]
174mod tests {
175    use super::*;
176
177    struct FailAfter(usize);
178
179    impl Write for FailAfter {
180        fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
181            if self.0 == 0 {
182                Err(io::Error::other("full"))
183            } else {
184                self.0 = self.0.saturating_sub(buf.len());
185                Ok(buf.len())
186            }
187        }
188
189        fn flush(&mut self) -> io::Result<()> {
190            Ok(())
191        }
192    }
193
194    #[derive(Default)]
195    struct Chunks {
196        writes: Vec<Vec<u8>>,
197    }
198
199    impl Write for Chunks {
200        fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
201            self.writes.push(buf.to_vec());
202            Ok(buf.len())
203        }
204
205        fn flush(&mut self) -> io::Result<()> {
206            Ok(())
207        }
208    }
209
210    impl Chunks {
211        fn joined(&self) -> Vec<u8> {
212            self.writes.iter().flatten().copied().collect()
213        }
214    }
215
216    #[test]
217    fn io_writer_buffers_and_flushes_in_order() {
218        let mut writer = IoWriter::with_capacity(8, Vec::new());
219        writer.write(b"abc");
220        writer.write_byte(b'd');
221        writer.write(b"efghij");
222        writer.write_with(4, |region| {
223            region.copy_from_slice(b"KLMN");
224            2
225        });
226        writer.write(b"0123456789abcdef");
227        assert_eq!(writer.finish().unwrap(), b"abcdefghijKL0123456789abcdef");
228    }
229
230    #[test]
231    fn io_writer_reports_the_first_error() {
232        let mut writer = IoWriter::with_capacity(4, FailAfter(6));
233        writer.write(b"abcdef");
234        writer.write(b"ghij");
235        writer.write(b"klmn");
236        assert!(writer.into_result().is_err());
237    }
238
239    #[test]
240    fn io_writer_keeps_writing_after_the_first_error() {
241        let mut writer = IoWriter::with_capacity(4, FailAfter(0));
242        for _ in 0..100 {
243            writer.write(b"abcdefgh");
244            writer.write_byte(b'x');
245        }
246        let error = writer.into_result().expect_err("error expected");
247        assert_eq!(error.to_string(), "full");
248    }
249
250    #[test]
251    fn io_writer_passes_large_writes_through() {
252        let mut writer = IoWriter::with_capacity(8, Chunks::default());
253        writer.write(b"ab");
254        writer.write(b"0123456789");
255        writer.write(b"cd");
256        let chunks = writer.finish().unwrap();
257        assert_eq!(chunks.writes.len(), 3);
258        assert_eq!(
259            chunks.writes.first().map(Vec::as_slice),
260            Some(b"ab".as_ref())
261        );
262        assert_eq!(
263            chunks.writes.get(1).map(Vec::as_slice),
264            Some(b"0123456789".as_ref())
265        );
266        assert_eq!(chunks.joined(), b"ab0123456789cd");
267    }
268
269    #[test]
270    fn io_writer_byte_writes_never_exceed_the_buffer() {
271        let mut writer = IoWriter::with_capacity(4, Chunks::default());
272        for byte in b"abcdefghij" {
273            writer.write_byte(*byte);
274        }
275        let chunks = writer.finish().unwrap();
276        assert!(
277            chunks.writes.iter().all(|chunk| chunk.len() <= 4),
278            "{:?}",
279            chunks.writes
280        );
281        assert_eq!(chunks.joined(), b"abcdefghij");
282    }
283
284    #[test]
285    fn io_writer_write_with_grows_beyond_the_buffer() {
286        let mut writer = IoWriter::with_capacity(4, Chunks::default());
287        writer.write(b"ab");
288        writer.write_with(16, |region| {
289            for (slot, byte) in region.iter_mut().zip(b"0123456789".iter()) {
290                *slot = *byte;
291            }
292            10
293        });
294        writer.write(b"cd");
295        let chunks = writer.finish().unwrap();
296        assert_eq!(chunks.joined(), b"ab0123456789cd");
297    }
298
299    #[test]
300    fn io_writer_write_with_keeps_only_the_filled_prefix() {
301        let mut writer = IoWriter::with_capacity(64, Vec::new());
302        writer.write_with(8, |region| {
303            assert_eq!(region.len(), 8);
304            assert!(region.iter().all(|byte| *byte == 0));
305            region.iter_mut().for_each(|slot| *slot = b'z');
306            3
307        });
308        writer.write_with(8, |_| 0);
309        writer.write(b"!");
310        assert_eq!(writer.finish().unwrap(), b"zzz!");
311    }
312
313    #[test]
314    fn io_writer_write_with_caps_the_reported_length() {
315        let mut writer = IoWriter::with_capacity(64, Vec::new());
316        writer.write_with(4, |region| {
317            region.copy_from_slice(b"abcd");
318            usize::MAX
319        });
320        assert_eq!(writer.finish().unwrap(), b"abcd");
321    }
322
323    #[test]
324    fn vec_write_with_keeps_only_the_filled_prefix() {
325        let mut out = b"x".to_vec();
326        out.write_with(6, |region| {
327            region[..3].copy_from_slice(b"abc");
328            3
329        });
330        assert_eq!(out, b"xabc");
331    }
332
333    #[test]
334    fn vec_write_with_zeroes_the_region() {
335        let mut out = Vec::new();
336        out.write_with(8, |region| {
337            region.iter_mut().for_each(|slot| *slot = b'a');
338            8
339        });
340        out.truncate(4);
341        out.write_with(8, |region| {
342            assert!(region.iter().all(|byte| *byte == 0));
343            0
344        });
345        assert_eq!(out, b"aaaa");
346    }
347
348    #[test]
349    fn default_write_with_matches_the_vec_implementation() {
350        struct Plain(Vec<u8>);
351        impl Writer for Plain {
352            fn write(&mut self, bytes: &[u8]) {
353                self.0.extend_from_slice(bytes);
354            }
355        }
356
357        let mut plain = Plain(Vec::new());
358        plain.write_byte(b'a');
359        plain.write_with(6, |region| {
360            region.iter_mut().for_each(|slot| *slot = b'b');
361            4
362        });
363        assert_eq!(plain.0, b"abbbb");
364    }
365}