Skip to main content

ax_io/utils/
copy.rs

1#[cfg(feature = "alloc")]
2use alloc::{collections::vec_deque::VecDeque, vec::Vec};
3use core::{io::BorrowedBuf, mem::MaybeUninit};
4
5use crate::{BufReader, BufWriter, DEFAULT_BUF_SIZE, Error, Read, Result, Write};
6
7/// Copies the entire contents of a reader into a writer.
8///
9/// This function will continuously read data from `reader` and then
10/// write it into `writer` in a streaming fashion until `reader`
11/// returns EOF.
12///
13/// On success, the total number of bytes that were copied from
14/// `reader` to `writer` is returned.
15///
16/// See [`std::io::copy`] for more details.
17pub fn copy<R, W>(reader: &mut R, writer: &mut W) -> Result<u64>
18where
19    R: Read + ?Sized,
20    W: Write + ?Sized,
21{
22    let read_buf = BufferedReaderSpec::buffer_size(reader);
23    let write_buf = BufferedWriterSpec::buffer_size(writer);
24
25    if read_buf >= DEFAULT_BUF_SIZE && read_buf >= write_buf {
26        return BufferedReaderSpec::copy_to(reader, writer);
27    }
28
29    BufferedWriterSpec::copy_from(writer, reader)
30}
31
32/// Fallback [`copy`] implementation using a stack-allocated buffer.
33pub fn stack_buffer_copy<R, W>(reader: &mut R, writer: &mut W) -> Result<u64>
34where
35    R: Read + ?Sized,
36    W: Write + ?Sized,
37{
38    let buf: &mut [_] = &mut [MaybeUninit::uninit(); DEFAULT_BUF_SIZE];
39    let mut buf: BorrowedBuf<'_, u8> = buf.into();
40
41    let mut len = 0;
42
43    loop {
44        match reader.read_buf(buf.unfilled()) {
45            Ok(()) => {}
46            Err(e) if e.canonicalize() == Error::Interrupted => continue,
47            Err(e) => return Err(e),
48        };
49
50        if buf.filled().is_empty() {
51            break;
52        }
53
54        len += buf.filled().len() as u64;
55        writer.write_all(buf.filled())?;
56        buf.clear();
57    }
58
59    Ok(len)
60}
61
62/// Specialization of the read-write loop that reuses the internal
63/// buffer of a BufReader. If there's no buffer then the writer side
64/// should be used instead.
65trait BufferedReaderSpec {
66    fn buffer_size(&self) -> usize;
67
68    fn copy_to(&mut self, to: &mut (impl Write + ?Sized)) -> Result<u64>;
69}
70
71impl<T> BufferedReaderSpec for T
72where
73    Self: Read,
74    T: ?Sized,
75{
76    #[inline]
77    default fn buffer_size(&self) -> usize {
78        0
79    }
80
81    default fn copy_to(&mut self, _to: &mut (impl Write + ?Sized)) -> Result<u64> {
82        unreachable!("only called from specializations")
83    }
84}
85
86impl BufferedReaderSpec for &[u8] {
87    fn buffer_size(&self) -> usize {
88        // prefer this specialization since the source "buffer" is all we'll ever need,
89        // even if it's small
90        usize::MAX
91    }
92
93    fn copy_to(&mut self, to: &mut (impl Write + ?Sized)) -> Result<u64> {
94        let len = self.len();
95        to.write_all(self)?;
96        *self = &self[len..];
97        Ok(len as u64)
98    }
99}
100
101#[cfg(feature = "alloc")]
102impl BufferedReaderSpec for VecDeque<u8> {
103    fn buffer_size(&self) -> usize {
104        // prefer this specialization since the source "buffer" is all we'll ever need,
105        // even if it's small
106        usize::MAX
107    }
108
109    fn copy_to(&mut self, to: &mut (impl Write + ?Sized)) -> Result<u64> {
110        let len = self.len();
111        let (front, back) = self.as_slices();
112        to.write_all(front)?;
113        to.write_all(back)?;
114        self.clear();
115        Ok(len as u64)
116    }
117}
118
119impl<I> BufferedReaderSpec for BufReader<I>
120where
121    Self: Read,
122    I: ?Sized,
123{
124    fn buffer_size(&self) -> usize {
125        self.capacity()
126    }
127
128    fn copy_to(&mut self, to: &mut (impl Write + ?Sized)) -> Result<u64> {
129        let mut len = 0;
130
131        loop {
132            // Hack: this relies on `impl Read for BufReader` always calling fill_buf
133            // if the buffer is empty, even for empty slices.
134            // It can't be called directly here since specialization prevents us
135            // from adding I: Read
136            match self.read(&mut []) {
137                Ok(_) => {}
138                Err(e) if e.canonicalize() == Error::Interrupted => continue,
139                Err(e) => return Err(e),
140            }
141            let buf = self.buffer();
142            if self.buffer().is_empty() {
143                return Ok(len);
144            }
145
146            // In case the writer side is a BufWriter then its write_all
147            // implements an optimization that passes through large
148            // buffers to the underlying writer. That code path is #[cold]
149            // but we're still avoiding redundant memcopies when doing
150            // a copy between buffered inputs and outputs.
151            to.write_all(buf)?;
152            len += buf.len() as u64;
153            self.discard_buffer();
154        }
155    }
156}
157
158/// Specialization of the read-write loop that either uses a stack buffer
159/// or reuses the internal buffer of a BufWriter
160trait BufferedWriterSpec: Write {
161    fn buffer_size(&self) -> usize;
162
163    fn copy_from<R: Read + ?Sized>(&mut self, reader: &mut R) -> Result<u64>;
164}
165
166impl<W: Write + ?Sized> BufferedWriterSpec for W {
167    #[inline]
168    default fn buffer_size(&self) -> usize {
169        0
170    }
171
172    default fn copy_from<R: Read + ?Sized>(&mut self, reader: &mut R) -> Result<u64> {
173        stack_buffer_copy(reader, self)
174    }
175}
176
177#[cfg(feature = "alloc")]
178impl BufferedWriterSpec for Vec<u8> {
179    fn buffer_size(&self) -> usize {
180        core::cmp::max(DEFAULT_BUF_SIZE, self.capacity() - self.len())
181    }
182
183    fn copy_from<R: Read + ?Sized>(&mut self, reader: &mut R) -> Result<u64> {
184        reader
185            .read_to_end(self)
186            .map(|bytes| u64::try_from(bytes).expect("usize overflowed u64"))
187    }
188}
189
190impl<I: Write + ?Sized> BufferedWriterSpec for BufWriter<I> {
191    fn buffer_size(&self) -> usize {
192        self.capacity()
193    }
194
195    fn copy_from<R: Read + ?Sized>(&mut self, reader: &mut R) -> Result<u64> {
196        if self.capacity() < DEFAULT_BUF_SIZE {
197            return stack_buffer_copy(reader, self);
198        }
199
200        let mut len = 0;
201        let mut init = false;
202
203        loop {
204            let buf = self.buffer_mut();
205            let mut read_buf: BorrowedBuf<'_, u8> = buf.spare_capacity_mut().into();
206
207            if init {
208                // SAFETY: init is either 0 or the init_len from the previous iteration.
209                unsafe { read_buf.set_init() };
210            }
211
212            if read_buf.capacity() >= DEFAULT_BUF_SIZE {
213                let mut cursor = read_buf.unfilled();
214                match reader.read_buf(cursor.reborrow()) {
215                    Ok(()) => {
216                        let bytes_read = cursor.written();
217
218                        if bytes_read == 0 {
219                            return Ok(len);
220                        }
221
222                        init = read_buf.is_init();
223                        len += bytes_read as u64;
224
225                        // SAFETY: BorrowedBuf guarantees all of its filled bytes are init
226                        unsafe { buf.set_len(buf.len() + bytes_read) };
227
228                        // Read again if the buffer still has enough capacity, as BufWriter itself
229                        // would do This will occur if the reader returns
230                        // short reads
231                    }
232                    Err(ref e) if e.canonicalize() == Error::Interrupted => {}
233                    Err(e) => return Err(e),
234                }
235            } else {
236                self.flush_buf()?;
237            }
238        }
239    }
240}
241
242#[cfg(test)]
243struct AxtestShortReader<'a> {
244    remaining: &'a [u8],
245    max_read: usize,
246    largest_request: usize,
247}
248
249#[cfg(test)]
250impl Read for AxtestShortReader<'_> {
251    fn read(&mut self, buf: &mut [u8]) -> Result<usize> {
252        self.largest_request = self.largest_request.max(buf.len());
253
254        let bytes = self.remaining.len().min(self.max_read).min(buf.len());
255        let (copied, remaining) = self.remaining.split_at(bytes);
256        buf[..bytes].copy_from_slice(copied);
257        self.remaining = remaining;
258
259        Ok(bytes)
260    }
261}
262
263#[cfg(test)]
264struct AxtestFixedWriter<'a> {
265    output: &'a mut [u8],
266    written: usize,
267    largest_write: usize,
268}
269
270#[cfg(test)]
271impl AxtestFixedWriter<'_> {
272    fn filled(&self) -> &[u8] {
273        &self.output[..self.written]
274    }
275}
276
277#[cfg(test)]
278impl Write for AxtestFixedWriter<'_> {
279    fn write(&mut self, buf: &[u8]) -> Result<usize> {
280        self.largest_write = self.largest_write.max(buf.len());
281
282        let bytes = (self.output.len() - self.written).min(buf.len());
283        let end = self.written + bytes;
284        self.output[self.written..end].copy_from_slice(&buf[..bytes]);
285        self.written = end;
286
287        Ok(bytes)
288    }
289
290    fn flush(&mut self) -> Result<()> {
291        Ok(())
292    }
293}
294
295#[cfg(test)]
296mod tests {
297    use super::*;
298
299    #[test]
300    fn copy_constants_hold() {
301        let source = *b"stack-buffer-copy";
302        let mut reader = AxtestShortReader {
303            remaining: &source,
304            max_read: 3,
305            largest_request: 0,
306        };
307        let mut output = [0; 17];
308        let mut writer = AxtestFixedWriter {
309            output: &mut output,
310            written: 0,
311            largest_write: 0,
312        };
313
314        assert_eq!(
315            stack_buffer_copy(&mut reader, &mut writer),
316            Ok(source.len() as u64)
317        );
318        assert_eq!(writer.filled(), source);
319        assert!(reader.remaining.is_empty());
320        assert_eq!(reader.largest_request, DEFAULT_BUF_SIZE);
321
322        let mut reader = AxtestShortReader {
323            remaining: &source,
324            max_read: source.len(),
325            largest_request: 0,
326        };
327        let mut short_output = [0; 4];
328        let mut short_writer = AxtestFixedWriter {
329            output: &mut short_output,
330            written: 0,
331            largest_write: 0,
332        };
333
334        assert_eq!(
335            stack_buffer_copy(&mut reader, &mut short_writer),
336            Err(Error::WriteZero)
337        );
338    }
339
340    #[test]
341    fn copy_buffered_reader_spec_hold() {
342        let source = [0x5a; DEFAULT_BUF_SIZE * 2 + 11];
343        let mut reader = BufReader::with_capacity(
344            DEFAULT_BUF_SIZE * 2,
345            AxtestShortReader {
346                remaining: &source,
347                max_read: source.len(),
348                largest_request: 0,
349            },
350        );
351        let mut output = [0; DEFAULT_BUF_SIZE * 2 + 11];
352        let mut writer = AxtestFixedWriter {
353            output: &mut output,
354            written: 0,
355            largest_write: 0,
356        };
357
358        assert_eq!(copy(&mut reader, &mut writer), Ok(source.len() as u64));
359        assert_eq!(writer.filled(), source);
360        assert_eq!(writer.largest_write, DEFAULT_BUF_SIZE * 2);
361        assert!(reader.into_inner().remaining.is_empty());
362    }
363
364    #[test]
365    fn copy_slice_specialization_hold() {
366        let source = [0x7b; DEFAULT_BUF_SIZE + 17];
367        let mut reader = source.as_slice();
368        let mut output = [0; DEFAULT_BUF_SIZE + 17];
369        let mut writer = AxtestFixedWriter {
370            output: &mut output,
371            written: 0,
372            largest_write: 0,
373        };
374
375        assert_eq!(copy(&mut reader, &mut writer), Ok(source.len() as u64));
376        assert_eq!(writer.filled(), source);
377        assert_eq!(writer.largest_write, source.len());
378        assert!(reader.is_empty());
379    }
380}