Skip to main content

uv_fs/
read.rs

1use std::io::{self, Read};
2
3const READ_BUFFER_SIZE: usize = 8 * 1024;
4
5/// A reader that validates its contents while reading.
6#[must_use]
7pub struct ValidatedReader<Reader> {
8    reader: Reader,
9    prefix: Option<String>,
10    require_utf8: bool,
11}
12
13impl<Reader> ValidatedReader<Reader> {
14    /// Create a new validated reader.
15    pub fn new(reader: Reader) -> Self {
16        Self {
17            reader,
18            prefix: None,
19            require_utf8: false,
20        }
21    }
22
23    /// Require the contents to start with `prefix`.
24    ///
25    /// If the prefix does not match, [`Self::read`] does not read any bytes after the prefix.
26    pub fn require_prefix(mut self, prefix: impl Into<String>) -> Self {
27        self.prefix = Some(prefix.into());
28        self
29    }
30
31    /// Require the contents to be valid UTF-8 text.
32    ///
33    /// The contents are read incrementally and rejected as soon as a NUL byte or invalid UTF-8 is
34    /// encountered.
35    pub fn require_utf8(mut self) -> Self {
36        self.require_utf8 = true;
37        self
38    }
39}
40
41impl<Reader: Read> ValidatedReader<Reader> {
42    /// Read the contents, returning `None` if a configured validation fails.
43    pub fn read(self) -> io::Result<Option<Vec<u8>>> {
44        let Self {
45            mut reader,
46            prefix,
47            require_utf8,
48        } = self;
49
50        let mut contents = if let Some(prefix) = prefix {
51            let mut contents = vec![0; prefix.len()];
52            match reader.read_exact(&mut contents) {
53                Ok(()) if contents == prefix.as_bytes() => {}
54                Ok(()) => return Ok(None),
55                Err(err) if err.kind() == io::ErrorKind::UnexpectedEof => return Ok(None),
56                Err(err) => return Err(err),
57            }
58            contents
59        } else {
60            Vec::new()
61        };
62
63        if !require_utf8 {
64            reader.read_to_end(&mut contents)?;
65            return Ok(Some(contents));
66        }
67
68        if contents.contains(&0) {
69            return Ok(None);
70        }
71
72        let mut valid_utf_8_len = contents.len();
73        let mut buffer = [0u8; READ_BUFFER_SIZE];
74        loop {
75            let count = match reader.read(&mut buffer) {
76                Ok(0) => break,
77                Ok(count) => count,
78                Err(err) if err.kind() == io::ErrorKind::Interrupted => continue,
79                Err(err) => return Err(err),
80            };
81
82            let chunk = &buffer[..count];
83            if chunk.contains(&0) {
84                return Ok(None);
85            }
86            contents.extend_from_slice(chunk);
87            match std::str::from_utf8(&contents[valid_utf_8_len..]) {
88                Ok(_) => valid_utf_8_len = contents.len(),
89                Err(err) if err.error_len().is_some() => return Ok(None),
90                Err(err) => valid_utf_8_len += err.valid_up_to(),
91            }
92        }
93
94        Ok((valid_utf_8_len == contents.len()).then_some(contents))
95    }
96}
97
98#[cfg(test)]
99mod tests {
100    use std::io::{self, Cursor, Read};
101
102    use super::ValidatedReader;
103
104    struct ErrorReader;
105
106    impl Read for ErrorReader {
107        fn read(&mut self, _buffer: &mut [u8]) -> io::Result<usize> {
108            Err(io::Error::other("read past rejected content"))
109        }
110    }
111
112    #[test]
113    fn stops_at_prefix_mismatch() -> io::Result<()> {
114        let reader = Cursor::new(b"# ").chain(ErrorReader);
115        assert!(
116            ValidatedReader::new(reader)
117                .require_prefix("#!")
118                .read()?
119                .is_none()
120        );
121        Ok(())
122    }
123
124    #[test]
125    fn reads_without_utf_8_validation() -> io::Result<()> {
126        assert_eq!(
127            ValidatedReader::new(Cursor::new(b"#!\xff"))
128                .require_prefix("#!")
129                .read()?,
130            Some(b"#!\xff".to_vec())
131        );
132        Ok(())
133    }
134
135    #[test]
136    fn stops_at_binary_content() -> io::Result<()> {
137        for marker in [0, 0xff] {
138            let reader = Cursor::new([b'#', b'!', marker]).chain(ErrorReader);
139            assert!(
140                ValidatedReader::new(reader)
141                    .require_prefix("#!")
142                    .require_utf8()
143                    .read()?
144                    .is_none()
145            );
146        }
147        Ok(())
148    }
149
150    #[test]
151    fn handles_split_utf_8() -> io::Result<()> {
152        let reader = Cursor::new(b"#!\xc3").chain(Cursor::new(b"\xa9"));
153        assert_eq!(
154            ValidatedReader::new(reader)
155                .require_prefix("#!")
156                .require_utf8()
157                .read()?,
158            Some(b"#!\xc3\xa9".to_vec())
159        );
160        Ok(())
161    }
162}