Skip to main content

concinnity_core/decode/
reader.rs

1//! Sequential bounds-checked reader over a byte buffer. Every accessor returns
2//! `Result`, so a decoder written against it cannot index past the end of its
3//! input no matter what lengths the input declares. Reading a fixed-width
4//! integer goes through `array`, which yields an owned `[u8; N]` and removes
5//! the `try_into().unwrap()` that hand-rolled cursors need.
6
7use alloc::format;
8use alloc::string::String;
9
10/// Reader over `bytes`, tracking a cursor and the payload name used in errors.
11#[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    /// A reader positioned at the start of `bytes`. `label` names the payload
20    /// kind in every error this reader produces.
21    pub fn new(bytes: &'a [u8], label: &'static str) -> Self {
22        Self {
23            bytes,
24            pos: 0,
25            label,
26        }
27    }
28
29    /// Open a tagged payload: prove it is long enough to hold a `header_bytes`
30    /// header and that it opens with `magic`, then position the reader just
31    /// past the magic. Lets a decoder report a short buffer once rather than
32    /// once per header field.
33    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    /// Current byte offset.
58    pub fn position(&self) -> usize {
59        self.pos
60    }
61
62    /// Bytes left between the cursor and the end of the buffer.
63    pub fn remaining(&self) -> usize {
64        self.bytes.len().saturating_sub(self.pos)
65    }
66
67    /// Whether no bytes remain.
68    pub fn is_empty(&self) -> bool {
69        self.remaining() == 0
70    }
71
72    /// Total buffer length, including bytes already consumed.
73    pub fn len(&self) -> usize {
74        self.bytes.len()
75    }
76
77    /// Consume `n` bytes and return them, or report where the buffer ran out.
78    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    /// Consume exactly `N` bytes as an owned array, ready for `from_le_bytes`.
99    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    /// Read one byte.
106    pub fn u8(&mut self) -> Result<u8, String> {
107        Ok(u8::from_le_bytes(self.array::<1>()?))
108    }
109
110    /// Read a little-endian `u16`.
111    pub fn u16(&mut self) -> Result<u16, String> {
112        Ok(u16::from_le_bytes(self.array::<2>()?))
113    }
114
115    /// Read a little-endian `u32`.
116    pub fn u32(&mut self) -> Result<u32, String> {
117        Ok(u32::from_le_bytes(self.array::<4>()?))
118    }
119
120    /// Read a little-endian `u64`.
121    pub fn u64(&mut self) -> Result<u64, String> {
122        Ok(u64::from_le_bytes(self.array::<8>()?))
123    }
124
125    /// Read a little-endian `i32`.
126    pub fn i32(&mut self) -> Result<i32, String> {
127        Ok(i32::from_le_bytes(self.array::<4>()?))
128    }
129
130    /// Read a little-endian `f32`.
131    pub fn f32(&mut self) -> Result<f32, String> {
132        Ok(f32::from_le_bytes(self.array::<4>()?))
133    }
134
135    /// Advance past `n` bytes without returning them.
136    pub fn skip(&mut self, n: usize) -> Result<(), String> {
137        self.take(n).map(|_| ())
138    }
139
140    /// Move the cursor to an absolute offset, which must lie within the buffer.
141    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    /// Whether the bytes at the cursor equal `magic`, without consuming them.
155    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    // Consume `magic`, or report that the buffer does not start with it.
163    #[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    /// Everything from the cursor to the end, leaving the cursor in place.
176    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    // A declared length near usize::MAX must not wrap the cursor arithmetic
230    // into a range that looks in-bounds.
231    #[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    // The header length guard has to run before the magic read, so an empty
344    // buffer reports the short header rather than an end-of-buffer error.
345    #[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}