1use std::io::{self, Write};
8
9pub 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 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
83pub 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 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}