Skip to main content

hns_encoding/
lib.rs

1#![doc = "Little-endian wire encoding with allocation bounds and complete-input checks."]
2
3use thiserror::Error;
4
5#[derive(Clone, Debug, Eq, Error, PartialEq)]
6pub enum DecodeError {
7    #[error("unexpected end of input at byte {offset}; needed {needed} more byte(s)")]
8    UnexpectedEnd { offset: usize, needed: usize },
9    #[error("length {actual} exceeds configured maximum {maximum}")]
10    LengthExceedsBound { actual: usize, maximum: usize },
11    #[error("{remaining} trailing byte(s) remain")]
12    TrailingBytes { remaining: usize },
13    #[error("invalid value for {field}: {reason}")]
14    InvalidValue {
15        field: &'static str,
16        reason: &'static str,
17    },
18}
19
20#[derive(Clone, Copy, Debug)]
21pub struct Decoder<'input> {
22    input: &'input [u8],
23    position: usize,
24}
25
26impl<'input> Decoder<'input> {
27    pub const fn new(input: &'input [u8]) -> Self {
28        Self { input, position: 0 }
29    }
30
31    pub const fn position(&self) -> usize {
32        self.position
33    }
34
35    pub const fn remaining(&self) -> usize {
36        self.input.len() - self.position
37    }
38
39    pub fn read_u8(&mut self) -> Result<u8, DecodeError> {
40        Ok(self.read_array::<1>()?[0])
41    }
42
43    pub fn read_u16_le(&mut self) -> Result<u16, DecodeError> {
44        Ok(u16::from_le_bytes(self.read_array()?))
45    }
46
47    pub fn read_u32_le(&mut self) -> Result<u32, DecodeError> {
48        Ok(u32::from_le_bytes(self.read_array()?))
49    }
50
51    pub fn read_u64_le(&mut self) -> Result<u64, DecodeError> {
52        Ok(u64::from_le_bytes(self.read_array()?))
53    }
54
55    pub fn read_compact_size(&mut self) -> Result<u64, DecodeError> {
56        match self.read_u8()? {
57            value @ 0x00..=0xfc => Ok(u64::from(value)),
58            0xfd => {
59                let value = u64::from(self.read_u16_le()?);
60                if value < 0xfd {
61                    return Err(DecodeError::InvalidValue {
62                        field: "compact size",
63                        reason: "noncanonical u16 encoding",
64                    });
65                }
66                Ok(value)
67            }
68            0xfe => {
69                let value = u64::from(self.read_u32_le()?);
70                if value <= u64::from(u16::MAX) {
71                    return Err(DecodeError::InvalidValue {
72                        field: "compact size",
73                        reason: "noncanonical u32 encoding",
74                    });
75                }
76                Ok(value)
77            }
78            0xff => {
79                let value = self.read_u64_le()?;
80                if value <= u64::from(u32::MAX) {
81                    return Err(DecodeError::InvalidValue {
82                        field: "compact size",
83                        reason: "noncanonical u64 encoding",
84                    });
85                }
86                Ok(value)
87            }
88        }
89    }
90
91    pub fn read_compact_usize(
92        &mut self,
93        maximum: usize,
94        _field: &'static str,
95    ) -> Result<usize, DecodeError> {
96        let value = self.read_compact_size()?;
97        let value = usize::try_from(value).map_err(|_| DecodeError::LengthExceedsBound {
98            actual: usize::MAX,
99            maximum,
100        })?;
101        if value > maximum {
102            return Err(DecodeError::LengthExceedsBound {
103                actual: value,
104                maximum,
105            });
106        }
107        Ok(value)
108    }
109
110    pub fn read_varbytes(
111        &mut self,
112        maximum: usize,
113        field: &'static str,
114    ) -> Result<Vec<u8>, DecodeError> {
115        let length = self.read_compact_usize(maximum, field)?;
116        self.read_bounded_vec(length, maximum)
117    }
118
119    pub fn read_array<const LENGTH: usize>(&mut self) -> Result<[u8; LENGTH], DecodeError> {
120        let bytes = self.read_slice(LENGTH)?;
121        let mut output = [0_u8; LENGTH];
122        output.copy_from_slice(bytes);
123        Ok(output)
124    }
125
126    pub fn read_slice(&mut self, length: usize) -> Result<&'input [u8], DecodeError> {
127        let end = self
128            .position
129            .checked_add(length)
130            .ok_or(DecodeError::LengthExceedsBound {
131                actual: usize::MAX,
132                maximum: self.remaining(),
133            })?;
134        if end > self.input.len() {
135            return Err(DecodeError::UnexpectedEnd {
136                offset: self.position,
137                needed: end - self.input.len(),
138            });
139        }
140        let bytes = &self.input[self.position..end];
141        self.position = end;
142        Ok(bytes)
143    }
144
145    pub fn read_bounded_vec(
146        &mut self,
147        length: usize,
148        maximum: usize,
149    ) -> Result<Vec<u8>, DecodeError> {
150        if length > maximum {
151            return Err(DecodeError::LengthExceedsBound {
152                actual: length,
153                maximum,
154            });
155        }
156        Ok(self.read_slice(length)?.to_vec())
157    }
158
159    pub fn finish(self) -> Result<(), DecodeError> {
160        if self.remaining() == 0 {
161            Ok(())
162        } else {
163            Err(DecodeError::TrailingBytes {
164                remaining: self.remaining(),
165            })
166        }
167    }
168}
169
170#[derive(Clone, Debug, Default, Eq, PartialEq)]
171pub struct Encoder {
172    bytes: Vec<u8>,
173}
174
175impl Encoder {
176    pub const fn new() -> Self {
177        Self { bytes: Vec::new() }
178    }
179
180    pub fn with_capacity(capacity: usize) -> Self {
181        Self {
182            bytes: Vec::with_capacity(capacity),
183        }
184    }
185
186    pub fn put_u8(&mut self, value: u8) {
187        self.bytes.push(value);
188    }
189
190    pub fn put_u16_le(&mut self, value: u16) {
191        self.bytes.extend_from_slice(&value.to_le_bytes());
192    }
193
194    pub fn put_u32_le(&mut self, value: u32) {
195        self.bytes.extend_from_slice(&value.to_le_bytes());
196    }
197
198    pub fn put_u64_le(&mut self, value: u64) {
199        self.bytes.extend_from_slice(&value.to_le_bytes());
200    }
201
202    pub fn put_compact_size(&mut self, value: u64) {
203        match value {
204            0x00..=0xfc => self.put_u8(value as u8),
205            0xfd..=0xffff => {
206                self.put_u8(0xfd);
207                self.put_u16_le(value as u16);
208            }
209            0x1_0000..=0xffff_ffff => {
210                self.put_u8(0xfe);
211                self.put_u32_le(value as u32);
212            }
213            _ => {
214                self.put_u8(0xff);
215                self.put_u64_le(value);
216            }
217        }
218    }
219
220    pub fn put_varbytes(&mut self, value: &[u8]) {
221        self.put_compact_size(value.len() as u64);
222        self.put_bytes(value);
223    }
224
225    pub fn put_bytes(&mut self, value: &[u8]) {
226        self.bytes.extend_from_slice(value);
227    }
228
229    pub fn into_bytes(self) -> Vec<u8> {
230        self.bytes
231    }
232}
233
234#[cfg(test)]
235mod tests {
236    use super::*;
237
238    #[test]
239    fn integers_round_trip_little_endian() {
240        let mut encoder = Encoder::new();
241        encoder.put_u8(1);
242        encoder.put_u16_le(0x0302);
243        encoder.put_u32_le(0x0706_0504);
244        encoder.put_u64_le(0x0f0e_0d0c_0b0a_0908);
245
246        let bytes = encoder.into_bytes();
247        assert_eq!(bytes, (1_u8..=15).collect::<Vec<_>>());
248
249        let mut decoder = Decoder::new(&bytes);
250        assert_eq!(decoder.read_u8(), Ok(1));
251        assert_eq!(decoder.read_u16_le(), Ok(0x0302));
252        assert_eq!(decoder.read_u32_le(), Ok(0x0706_0504));
253        assert_eq!(decoder.read_u64_le(), Ok(0x0f0e_0d0c_0b0a_0908));
254        assert_eq!(decoder.finish(), Ok(()));
255    }
256
257    #[test]
258    fn rejects_truncation_trailing_bytes_and_oversized_allocation() {
259        let mut truncated = Decoder::new(&[1, 2, 3]);
260        assert!(matches!(
261            truncated.read_u32_le(),
262            Err(DecodeError::UnexpectedEnd { .. })
263        ));
264
265        let mut trailing = Decoder::new(&[1, 2]);
266        assert_eq!(trailing.read_u8(), Ok(1));
267        assert_eq!(
268            trailing.finish(),
269            Err(DecodeError::TrailingBytes { remaining: 1 })
270        );
271
272        let mut bounded = Decoder::new(&[0; 4]);
273        assert_eq!(
274            bounded.read_bounded_vec(4, 3),
275            Err(DecodeError::LengthExceedsBound {
276                actual: 4,
277                maximum: 3
278            })
279        );
280        assert_eq!(bounded.position(), 0);
281    }
282
283    #[test]
284    fn compact_sizes_are_minimal_and_bounded() {
285        for value in [
286            0,
287            0xfc,
288            0xfd,
289            u64::from(u16::MAX),
290            u64::from(u16::MAX) + 1,
291            u64::from(u32::MAX),
292            u64::from(u32::MAX) + 1,
293            u64::MAX,
294        ] {
295            let mut encoder = Encoder::new();
296            encoder.put_compact_size(value);
297            let bytes = encoder.into_bytes();
298            let mut decoder = Decoder::new(&bytes);
299            assert_eq!(decoder.read_compact_size(), Ok(value));
300            assert_eq!(decoder.finish(), Ok(()));
301        }
302        assert!(Decoder::new(&[0xfd, 0xfc, 0]).read_compact_size().is_err());
303        assert!(Decoder::new(&[4]).read_compact_usize(3, "items").is_err());
304    }
305}