1use std::io::{self, Read};
2
3const READ_BUFFER_SIZE: usize = 8 * 1024;
4
5#[must_use]
7pub struct ValidatedReader<Reader> {
8 reader: Reader,
9 prefix: Option<String>,
10 require_utf8: bool,
11}
12
13impl<Reader> ValidatedReader<Reader> {
14 pub fn new(reader: Reader) -> Self {
16 Self {
17 reader,
18 prefix: None,
19 require_utf8: false,
20 }
21 }
22
23 pub fn require_prefix(mut self, prefix: impl Into<String>) -> Self {
27 self.prefix = Some(prefix.into());
28 self
29 }
30
31 pub fn require_utf8(mut self) -> Self {
36 self.require_utf8 = true;
37 self
38 }
39}
40
41impl<Reader: Read> ValidatedReader<Reader> {
42 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}