1use alloc::format;
8use alloc::string::String;
9
10#[derive(Debug, Clone)]
12pub struct ByteReader<'a> {
13 bytes: &'a [u8],
14 pos: usize,
15 label: &'static str,
16}
17
18impl<'a> ByteReader<'a> {
19 pub fn new(bytes: &'a [u8], label: &'static str) -> Self {
22 Self {
23 bytes,
24 pos: 0,
25 label,
26 }
27 }
28
29 pub fn open_payload(
34 bytes: &'a [u8],
35 magic: u32,
36 header_bytes: usize,
37 label: &'static str,
38 ) -> Result<Self, String> {
39 if bytes.len() < header_bytes {
40 return Err(format!(
41 "{} payload too short: {} bytes (need at least {} for header)",
42 label,
43 bytes.len(),
44 header_bytes
45 ));
46 }
47 let mut r = Self::new(bytes, label);
48 let found = r.u32()?;
49 if found != magic {
50 return Err(format!(
51 "{label} payload magic 0x{found:08x} does not match expected 0x{magic:08x}"
52 ));
53 }
54 Ok(r)
55 }
56
57 pub fn position(&self) -> usize {
59 self.pos
60 }
61
62 pub fn remaining(&self) -> usize {
64 self.bytes.len().saturating_sub(self.pos)
65 }
66
67 pub fn is_empty(&self) -> bool {
69 self.remaining() == 0
70 }
71
72 pub fn len(&self) -> usize {
74 self.bytes.len()
75 }
76
77 pub fn take(&mut self, n: usize) -> Result<&'a [u8], String> {
79 let end = self.pos.checked_add(n).ok_or_else(|| {
80 format!(
81 "{} length overflow reading {} bytes at offset {}",
82 self.label, n, self.pos
83 )
84 })?;
85 let out = self.bytes.get(self.pos..end).ok_or_else(|| {
86 format!(
87 "unexpected end of {}: need {} bytes at offset {}, have {}",
88 self.label,
89 n,
90 self.pos,
91 self.bytes.len()
92 )
93 })?;
94 self.pos = end;
95 Ok(out)
96 }
97
98 pub fn array<const N: usize>(&mut self) -> Result<[u8; N], String> {
100 let mut out = [0u8; N];
101 out.copy_from_slice(self.take(N)?);
102 Ok(out)
103 }
104
105 pub fn u8(&mut self) -> Result<u8, String> {
107 Ok(u8::from_le_bytes(self.array::<1>()?))
108 }
109
110 pub fn u16(&mut self) -> Result<u16, String> {
112 Ok(u16::from_le_bytes(self.array::<2>()?))
113 }
114
115 pub fn u32(&mut self) -> Result<u32, String> {
117 Ok(u32::from_le_bytes(self.array::<4>()?))
118 }
119
120 pub fn u64(&mut self) -> Result<u64, String> {
122 Ok(u64::from_le_bytes(self.array::<8>()?))
123 }
124
125 pub fn i32(&mut self) -> Result<i32, String> {
127 Ok(i32::from_le_bytes(self.array::<4>()?))
128 }
129
130 pub fn f32(&mut self) -> Result<f32, String> {
132 Ok(f32::from_le_bytes(self.array::<4>()?))
133 }
134
135 pub fn skip(&mut self, n: usize) -> Result<(), String> {
137 self.take(n).map(|_| ())
138 }
139
140 pub fn seek(&mut self, pos: usize) -> Result<(), String> {
142 if pos > self.bytes.len() {
143 return Err(format!(
144 "{} seek to offset {} past end of {} bytes",
145 self.label,
146 pos,
147 self.bytes.len()
148 ));
149 }
150 self.pos = pos;
151 Ok(())
152 }
153
154 pub fn peek(&self, magic: &[u8]) -> bool {
156 self.pos
157 .checked_add(magic.len())
158 .and_then(|end| self.bytes.get(self.pos..end))
159 .is_some_and(|b| b == magic)
160 }
161
162 #[cfg(test)]
164 pub(crate) fn expect_magic(&mut self, magic: &[u8]) -> Result<(), String> {
165 let found = self.take(magic.len())?;
166 if found != magic {
167 return Err(format!(
168 "{} magic {:02x?} does not match expected {:02x?}",
169 self.label, found, magic
170 ));
171 }
172 Ok(())
173 }
174
175 pub fn remainder(&self) -> &'a [u8] {
177 self.bytes.get(self.pos..).unwrap_or(&[])
178 }
179}
180
181#[cfg(test)]
182mod tests {
183 use super::*;
184 use alloc::vec::Vec;
185
186 fn reader(bytes: &[u8]) -> ByteReader<'_> {
187 ByteReader::new(bytes, "test")
188 }
189
190 #[test]
191 fn reads_fixed_width_integers_in_order() {
192 let mut buf = Vec::new();
193 buf.extend_from_slice(&7u32.to_le_bytes());
194 buf.extend_from_slice(&9u16.to_le_bytes());
195 buf.extend_from_slice(&1.5f32.to_le_bytes());
196 buf.extend_from_slice(&(-3i32).to_le_bytes());
197 buf.extend_from_slice(&11u64.to_le_bytes());
198 buf.push(200);
199
200 let mut r = reader(&buf);
201 assert_eq!(r.u32().unwrap(), 7);
202 assert_eq!(r.u16().unwrap(), 9);
203 assert_eq!(r.f32().unwrap(), 1.5);
204 assert_eq!(r.i32().unwrap(), -3);
205 assert_eq!(r.u64().unwrap(), 11);
206 assert_eq!(r.u8().unwrap(), 200);
207 assert!(r.is_empty());
208 }
209
210 #[test]
211 fn take_advances_and_tracks_position() {
212 let buf = [1u8, 2, 3, 4, 5];
213 let mut r = reader(&buf);
214 assert_eq!(r.take(2).unwrap(), &[1, 2]);
215 assert_eq!(r.position(), 2);
216 assert_eq!(r.remaining(), 3);
217 assert_eq!(r.len(), 5);
218 }
219
220 #[test]
221 fn take_past_end_errors_instead_of_panicking() {
222 let buf = [1u8, 2, 3];
223 let mut r = reader(&buf);
224 let err = r.take(4).unwrap_err();
225 assert!(err.contains("unexpected end of test"), "{}", err);
226 assert!(err.contains("have 3"), "{}", err);
227 }
228
229 #[test]
232 fn take_length_overflow_errors() {
233 let buf = [1u8, 2, 3, 4];
234 let mut r = reader(&buf);
235 r.skip(2).unwrap();
236 let err = r.take(usize::MAX).unwrap_err();
237 assert!(err.contains("length overflow"), "{}", err);
238 }
239
240 #[test]
241 fn failed_take_leaves_cursor_untouched() {
242 let buf = [1u8, 2, 3];
243 let mut r = reader(&buf);
244 r.skip(1).unwrap();
245 assert!(r.take(99).is_err());
246 assert_eq!(r.position(), 1);
247 assert_eq!(r.u8().unwrap(), 2);
248 }
249
250 #[test]
251 fn truncated_integer_read_errors() {
252 let buf = [1u8, 2];
253 let mut r = reader(&buf);
254 assert!(r.u32().is_err());
255 }
256
257 #[test]
258 fn seek_moves_cursor_and_rejects_past_end() {
259 let buf = [1u8, 2, 3, 4];
260 let mut r = reader(&buf);
261 r.seek(3).unwrap();
262 assert_eq!(r.u8().unwrap(), 4);
263 r.seek(4).unwrap();
264 assert!(r.is_empty());
265 assert!(r.seek(5).is_err());
266 }
267
268 #[test]
269 fn peek_does_not_consume() {
270 let buf = *b"CNB\0rest";
271 let mut r = reader(&buf);
272 assert!(r.peek(b"CNB\0"));
273 assert!(!r.peek(b"XXXX"));
274 assert_eq!(r.position(), 0);
275 r.expect_magic(b"CNB\0").unwrap();
276 assert_eq!(r.position(), 4);
277 }
278
279 #[test]
280 fn peek_past_end_is_false_not_a_panic() {
281 let buf = [1u8, 2];
282 let r = reader(&buf);
283 assert!(!r.peek(b"CNB\0"));
284 }
285
286 #[test]
287 fn expect_magic_reports_mismatch() {
288 let buf = *b"XXXXrest";
289 let mut r = reader(&buf);
290 let err = r.expect_magic(b"CNB\0").unwrap_err();
291 assert!(err.contains("does not match"), "{}", err);
292 }
293
294 #[test]
295 fn expect_magic_on_short_buffer_errors() {
296 let buf = *b"CN";
297 let mut r = reader(&buf);
298 assert!(r.expect_magic(b"CNB\0").is_err());
299 }
300
301 #[test]
302 fn remainder_returns_unconsumed_tail() {
303 let buf = [1u8, 2, 3, 4];
304 let mut r = reader(&buf);
305 r.skip(2).unwrap();
306 assert_eq!(r.remainder(), &[3, 4]);
307 assert_eq!(r.position(), 2);
308 }
309
310 const MAGIC: u32 = u32::from_le_bytes(*b"TEST");
311
312 fn tagged(fields: &[u32]) -> Vec<u8> {
313 let mut buf = MAGIC.to_le_bytes().to_vec();
314 for f in fields {
315 buf.extend_from_slice(&f.to_le_bytes());
316 }
317 buf
318 }
319
320 #[test]
321 fn open_payload_positions_past_the_magic() {
322 let bytes = tagged(&[7, 8]);
323 let mut r = ByteReader::open_payload(&bytes, MAGIC, 12, "test").unwrap();
324 assert_eq!(r.position(), 4);
325 assert_eq!(r.u32().unwrap(), 7);
326 assert_eq!(r.u32().unwrap(), 8);
327 }
328
329 #[test]
330 fn open_payload_rejects_a_short_header() {
331 let bytes = tagged(&[7]);
332 let err = ByteReader::open_payload(&bytes, MAGIC, 12, "test").unwrap_err();
333 assert!(err.contains("too short"), "{}", err);
334 }
335
336 #[test]
337 fn open_payload_rejects_a_wrong_magic() {
338 let bytes = tagged(&[7, 8]);
339 let err = ByteReader::open_payload(&bytes, 0xDEAD_BEEF, 12, "test").unwrap_err();
340 assert!(err.contains("magic"), "{}", err);
341 }
342
343 #[test]
346 fn open_payload_on_an_empty_buffer_reports_a_short_header() {
347 let err = ByteReader::open_payload(&[], MAGIC, 12, "test").unwrap_err();
348 assert!(err.contains("too short"), "{}", err);
349 }
350
351 #[test]
352 fn empty_buffer_reads_error() {
353 let mut r = reader(&[]);
354 assert!(r.is_empty());
355 assert_eq!(r.remainder(), &[] as &[u8]);
356 assert!(r.u8().is_err());
357 }
358}