Skip to main content

ironwork_rt/module/
codec.rs

1//! The codec traits, the writer and reader they use, the §4.1 and §4.4 impls, and the §4.6 macros.
2
3use std::collections::BTreeMap;
4
5use super::ModuleError;
6use super::leb::{self, LebError};
7use super::strings::{Interner, StringTable};
8
9pub trait Encode {
10    fn encode(&self, w: &mut Writer);
11}
12
13pub trait Decode: Sized {
14    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError>;
15}
16
17/// Section bytes being built, and the string table they intern into.
18#[derive(Debug, Default)]
19pub struct Writer {
20    bytes: Vec<u8>,
21    strings: Interner,
22}
23
24impl Writer {
25    pub fn new() -> Self {
26        Self::default()
27    }
28
29    pub fn byte(&mut self, byte: u8) {
30        self.bytes.push(byte);
31    }
32
33    pub fn bytes(&mut self, bytes: &[u8]) {
34        self.bytes.extend_from_slice(bytes);
35    }
36
37    pub fn leb(&mut self, value: u64) {
38        leb::write(&mut self.bytes, value);
39    }
40
41    pub fn zigzag(&mut self, value: i64) {
42        self.leb(leb::zigzag(value));
43    }
44
45    pub fn count(&mut self, count: usize) {
46        self.leb(count as u64);
47    }
48
49    /// Writes the string's index in the table, adding it on first use.
50    pub fn string(&mut self, text: &str) {
51        let index = self.strings.intern(text);
52        self.count(index);
53    }
54
55    /// The bytes written since the last `take`. The string table carries on.
56    pub fn take(&mut self) -> Vec<u8> {
57        std::mem::take(&mut self.bytes)
58    }
59
60    pub fn strings(&self) -> &StringTable {
61        self.strings.table()
62    }
63}
64
65/// A bounds-checked cursor over one section's bytes.
66#[derive(Clone, Debug)]
67pub struct Reader<'r> {
68    bytes: &'r [u8],
69    at: usize,
70    section: &'static str,
71    strings: &'r StringTable,
72}
73
74impl<'r> Reader<'r> {
75    pub fn new(section: &'static str, bytes: &'r [u8], strings: &'r StringTable) -> Self {
76        Self { bytes, at: 0, section, strings }
77    }
78
79    pub fn position(&self) -> usize {
80        self.at
81    }
82
83    pub fn remaining(&self) -> usize {
84        self.bytes.len() - self.at
85    }
86
87    pub fn malformed(&self, at: usize, reason: impl Into<String>) -> ModuleError {
88        ModuleError::Malformed { section: self.section, offset: at, reason: reason.into() }
89    }
90
91    pub fn byte(&mut self) -> Result<u8, ModuleError> {
92        Ok(self.bytes(1)?[0])
93    }
94
95    pub fn bytes(&mut self, len: usize) -> Result<&'r [u8], ModuleError> {
96        let taken = self.at.checked_add(len).and_then(|end| self.bytes.get(self.at..end));
97        let taken = taken.ok_or_else(|| {
98            self.malformed(self.at, format!("reads past the end: {len} wanted, {} left", self.remaining()))
99        })?;
100        self.at += len;
101        Ok(taken)
102    }
103
104    pub fn leb(&mut self) -> Result<u64, ModuleError> {
105        let (value, len) = leb::read(&self.bytes[self.at..]).map_err(|e| {
106            self.malformed(
107                self.at,
108                match e {
109                    LebError::End => "the bytes end inside an integer",
110                    LebError::OverLong => "an integer is not in its shortest form",
111                    LebError::Overflow => "an integer overflows 64 bits",
112                },
113            )
114        })?;
115        self.at += len;
116        Ok(value)
117    }
118
119    pub fn zigzag(&mut self) -> Result<i64, ModuleError> {
120        Ok(leb::unzigzag(self.leb()?))
121    }
122
123    /// A count of elements, each of which takes at least one byte, so no more than remain.
124    pub fn count(&mut self) -> Result<usize, ModuleError> {
125        let at = self.at;
126        let count = self.leb()?;
127        let remaining = self.remaining();
128        usize::try_from(count)
129            .ok()
130            .filter(|&n| n <= remaining)
131            .ok_or_else(|| self.malformed(at, format!("a count of {count} with {remaining} bytes left")))
132    }
133
134    pub fn string(&mut self) -> Result<&'r str, ModuleError> {
135        let at = self.at;
136        let index = self.leb()?;
137        let strings = self.strings;
138        usize::try_from(index)
139            .ok()
140            .and_then(|i| strings.get(i))
141            .ok_or_else(|| self.malformed(at, format!("string {index} of a table of {}", strings.len())))
142    }
143
144    /// Refuses bytes left after the last value.
145    pub fn finish(self) -> Result<(), ModuleError> {
146        match self.remaining() {
147            0 => Ok(()),
148            left => Err(self.malformed(self.at, format!("bytes left after the last value: {left}"))),
149        }
150    }
151}
152
153/// One `T` that fills `bytes` exactly.
154pub fn decode_all<T: Decode>(section: &'static str, bytes: &[u8], strings: &StringTable) -> Result<T, ModuleError> {
155    let mut r = Reader::new(section, bytes, strings);
156    let value = T::decode(&mut r)?;
157    r.finish()?;
158    Ok(value)
159}
160
161/// The `check` of a `codec_struct!` that names none.
162pub fn unchecked<T>(_: &T) -> Result<(), String> {
163    Ok(())
164}
165
166impl Encode for u8 {
167    fn encode(&self, w: &mut Writer) {
168        w.byte(*self);
169    }
170}
171
172impl Decode for u8 {
173    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
174        r.byte()
175    }
176}
177
178impl Encode for bool {
179    fn encode(&self, w: &mut Writer) {
180        w.byte(u8::from(*self));
181    }
182}
183
184impl Decode for bool {
185    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
186        let at = r.position();
187        match r.byte()? {
188            0 => Ok(false),
189            1 => Ok(true),
190            other => Err(r.malformed(at, format!("bool {other}"))),
191        }
192    }
193}
194
195macro_rules! unsigned {
196    ($($t:ty),*) => {$(
197        impl Encode for $t {
198            fn encode(&self, w: &mut Writer) {
199                w.leb(*self as u64);
200            }
201        }
202
203        impl Decode for $t {
204            fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
205                let at = r.position();
206                let value = r.leb()?;
207                <$t>::try_from(value).map_err(|_| r.malformed(at, format!("{value} overflows {}", stringify!($t))))
208            }
209        }
210    )*};
211}
212
213unsigned!(u16, u32, u64, usize);
214
215macro_rules! signed {
216    ($($t:ty),*) => {$(
217        impl Encode for $t {
218            fn encode(&self, w: &mut Writer) {
219                w.zigzag(i64::from(*self));
220            }
221        }
222
223        impl Decode for $t {
224            fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
225                let at = r.position();
226                let value = r.zigzag()?;
227                <$t>::try_from(value).map_err(|_| r.malformed(at, format!("{value} overflows {}", stringify!($t))))
228            }
229        }
230    )*};
231}
232
233signed!(i16, i32, i64);
234
235impl Encode for char {
236    fn encode(&self, w: &mut Writer) {
237        w.leb(u64::from(*self));
238    }
239}
240
241impl Decode for char {
242    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
243        let at = r.position();
244        let value = r.leb()?;
245        u32::try_from(value)
246            .ok()
247            .and_then(char::from_u32)
248            .ok_or_else(|| r.malformed(at, format!("{value:#x} is not a Unicode scalar value")))
249    }
250}
251
252const CANONICAL_NAN: u64 = 0x7FF8_0000_0000_0000;
253
254impl Encode for f64 {
255    fn encode(&self, w: &mut Writer) {
256        let bits = if self.is_nan() { CANONICAL_NAN } else { self.to_bits() };
257        w.bytes(&bits.to_le_bytes());
258    }
259}
260
261impl Decode for f64 {
262    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
263        let at = r.position();
264        let mut bits = [0u8; 8];
265        bits.copy_from_slice(r.bytes(8)?);
266        let value = f64::from_le_bytes(bits);
267        if value.is_nan() && value.to_bits() != CANONICAL_NAN {
268            return Err(r.malformed(at, "a NaN other than the canonical one"));
269        }
270        Ok(value)
271    }
272}
273
274impl Encode for String {
275    fn encode(&self, w: &mut Writer) {
276        w.string(self);
277    }
278}
279
280impl Decode for String {
281    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
282        r.string().map(str::to_owned)
283    }
284}
285
286impl<T: Encode> Encode for Vec<T> {
287    fn encode(&self, w: &mut Writer) {
288        w.count(self.len());
289        for item in self {
290            item.encode(w);
291        }
292    }
293}
294
295impl<T: Decode> Decode for Vec<T> {
296    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
297        let count = r.count()?;
298        let mut items = Vec::with_capacity(count);
299        for _ in 0..count {
300            items.push(T::decode(r)?);
301        }
302        Ok(items)
303    }
304}
305
306impl<T: Encode, const N: usize> Encode for [T; N] {
307    fn encode(&self, w: &mut Writer) {
308        for item in self {
309            item.encode(w);
310        }
311    }
312}
313
314impl<T: Decode, const N: usize> Decode for [T; N] {
315    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
316        let at = r.position();
317        if N > r.remaining() {
318            return Err(r.malformed(at, format!("an array of {N} with {} bytes left", r.remaining())));
319        }
320        let mut items = Vec::with_capacity(N);
321        for _ in 0..N {
322            items.push(T::decode(r)?);
323        }
324        items.try_into().map_err(|_| r.malformed(at, format!("an array of {N}")))
325    }
326}
327
328impl<T: Encode> Encode for Option<T> {
329    fn encode(&self, w: &mut Writer) {
330        match self {
331            None => w.byte(0),
332            Some(value) => {
333                w.byte(1);
334                value.encode(w);
335            }
336        }
337    }
338}
339
340impl<T: Decode> Decode for Option<T> {
341    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
342        let at = r.position();
343        match r.byte()? {
344            0 => Ok(None),
345            1 => T::decode(r).map(Some),
346            other => Err(r.malformed(at, format!("Option tag {other}"))),
347        }
348    }
349}
350
351impl<T: Encode, E: Encode> Encode for Result<T, E> {
352    fn encode(&self, w: &mut Writer) {
353        match self {
354            Ok(value) => {
355                w.byte(0);
356                value.encode(w);
357            }
358            Err(error) => {
359                w.byte(1);
360                error.encode(w);
361            }
362        }
363    }
364}
365
366impl<T: Decode, E: Decode> Decode for Result<T, E> {
367    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
368        let at = r.position();
369        match r.byte()? {
370            0 => T::decode(r).map(Ok),
371            1 => E::decode(r).map(Err),
372            other => Err(r.malformed(at, format!("Result tag {other}"))),
373        }
374    }
375}
376
377impl<T: Encode> Encode for Box<T> {
378    fn encode(&self, w: &mut Writer) {
379        (**self).encode(w);
380    }
381}
382
383impl<T: Decode> Decode for Box<T> {
384    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
385        T::decode(r).map(Box::new)
386    }
387}
388
389macro_rules! tuple {
390    ($($name:ident $index:tt),+) => {
391        impl<$($name: Encode),+> Encode for ($($name,)+) {
392            fn encode(&self, w: &mut Writer) {
393                $(self.$index.encode(w);)+
394            }
395        }
396
397        impl<$($name: Decode),+> Decode for ($($name,)+) {
398            fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
399                Ok(($($name::decode(r)?,)+))
400            }
401        }
402    };
403}
404
405tuple!(A 0, B 1);
406tuple!(A 0, B 1, C 2);
407tuple!(A 0, B 1, C 2, D 3);
408
409impl<K: Encode, V: Encode> Encode for BTreeMap<K, V> {
410    fn encode(&self, w: &mut Writer) {
411        w.count(self.len());
412        for (key, value) in self {
413            key.encode(w);
414            value.encode(w);
415        }
416    }
417}
418
419impl<K: Decode + Ord, V: Decode> Decode for BTreeMap<K, V> {
420    fn decode(r: &mut Reader<'_>) -> Result<Self, ModuleError> {
421        let count = r.count()?;
422        let mut map = BTreeMap::new();
423        for _ in 0..count {
424            let at = r.position();
425            let key = K::decode(r)?;
426            if map.last_key_value().is_some_and(|(last, _)| *last >= key) {
427                return Err(r.malformed(at, "map keys that do not strictly ascend"));
428            }
429            map.insert(key, V::decode(r)?);
430        }
431        Ok(map)
432    }
433}
434
435/// `Encode` and `Decode` from one field list; `check` names a `fn(&T) -> Result<(), String>`.
436#[macro_export]
437macro_rules! codec_struct {
438    ($ty:ident { $($field:ident),* $(,)? }) => {
439        $crate::codec_struct!($ty { $($field),* } check $crate::module::codec::unchecked);
440    };
441    ($ty:ident { $($field:ident),* $(,)? } check $check:path) => {
442        impl $crate::module::codec::Encode for $ty {
443            fn encode(&self, w: &mut $crate::module::codec::Writer) {
444                let $ty { $($field),* } = self;
445                $($crate::module::codec::Encode::encode($field, w);)*
446            }
447        }
448
449        impl $crate::module::codec::Decode for $ty {
450            fn decode(
451                r: &mut $crate::module::codec::Reader<'_>,
452            ) -> ::core::result::Result<Self, $crate::module::ModuleError> {
453                let at = r.position();
454                $(let $field = $crate::module::codec::Decode::decode(r)?;)*
455                let value = $ty { $($field),* };
456                $check(&value).map_err(|reason| r.malformed(at, reason))?;
457                ::core::result::Result::Ok(value)
458            }
459        }
460    };
461}
462
463/// `Encode` and `Decode` from the variants and their tags: `Unit = 0`, `Named { a } = 1`, `Tuple(a) = 2`.
464#[macro_export]
465macro_rules! codec_enum {
466    ($ty:ident {
467        $($variant:ident $({ $($field:ident),* $(,)? })? $(( $($elem:ident),* $(,)? ))? = $tag:literal),* $(,)?
468    }) => {
469        impl $crate::module::codec::Encode for $ty {
470            fn encode(&self, w: &mut $crate::module::codec::Writer) {
471                match self {
472                    $($ty::$variant $({ $($field),* })? $(( $($elem),* ))? => {
473                        w.leb($tag);
474                        $($($crate::module::codec::Encode::encode($field, w);)*)?
475                        $($($crate::module::codec::Encode::encode($elem, w);)*)?
476                    })*
477                }
478            }
479        }
480
481        impl $crate::module::codec::Decode for $ty {
482            fn decode(
483                r: &mut $crate::module::codec::Reader<'_>,
484            ) -> ::core::result::Result<Self, $crate::module::ModuleError> {
485                let at = r.position();
486                match r.leb()? {
487                    $($tag => {
488                        $($(let $field = $crate::module::codec::Decode::decode(r)?;)*)?
489                        $($(let $elem = $crate::module::codec::Decode::decode(r)?;)*)?
490                        ::core::result::Result::Ok($ty::$variant $({ $($field),* })? $(( $($elem),* ))?)
491                    })*
492                    tag => ::core::result::Result::Err(r.malformed(at, format!("{} has no tag {tag}", stringify!($ty)))),
493                }
494            }
495        }
496    };
497}
498
499#[cfg(test)]
500mod tests {
501    use super::*;
502
503    fn encoded<T: Encode>(value: &T) -> (Vec<u8>, StringTable) {
504        let mut w = Writer::new();
505        value.encode(&mut w);
506        (w.take(), w.strings().clone())
507    }
508
509    fn round_trip<T: Encode + Decode + PartialEq + std::fmt::Debug>(value: T) {
510        let (bytes, strings) = encoded(&value);
511        assert_eq!(decode_all::<T>("TEST", &bytes, &strings), Ok(value));
512    }
513
514    fn refused<T: Decode + std::fmt::Debug>(bytes: &[u8]) -> String {
515        match decode_all::<T>("TEST", bytes, &StringTable::default()) {
516            Err(ModuleError::Malformed { section: "TEST", reason, .. }) => reason,
517            other => panic!("{bytes:02X?} decoded as {other:?}"),
518        }
519    }
520
521    #[test]
522    fn integers_round_trip_at_their_limits() {
523        for v in [0, 1, 0x7F, 0x80, u8::MAX] {
524            round_trip(v);
525        }
526        for v in [0, 0x80, u16::MAX] {
527            round_trip(v);
528        }
529        for v in [0, 0x3FFF, 0x4000, u32::MAX] {
530            round_trip(v);
531        }
532        for v in [0, u64::from(u32::MAX) + 1, u64::MAX] {
533            round_trip(v);
534        }
535        for v in [0, usize::MAX] {
536            round_trip(v);
537        }
538        for v in [i16::MIN, -1, 0, 1, i16::MAX] {
539            round_trip(v);
540        }
541        for v in [i32::MIN, -64, 63, i32::MAX] {
542            round_trip(v);
543        }
544        for v in [i64::MIN, -1, 0, i64::MAX] {
545            round_trip(v);
546        }
547        round_trip(false);
548        round_trip(true);
549        for c in ['\0', 'A', 'é', '€', '\u{10FFFF}'] {
550            round_trip(c);
551        }
552    }
553
554    #[test]
555    fn integers_have_the_documented_bytes() {
556        assert_eq!(encoded(&1140u16).0, [0xF4, 0x08]);
557        assert_eq!(encoded(&-1i32).0, [0x01]);
558        assert_eq!(encoded(&-65i64).0, [0x81, 0x01]);
559        assert_eq!(encoded(&200u8).0, [200]);
560        assert_eq!(encoded(&true).0, [1]);
561        assert_eq!(encoded(&'€').0, [0xAC, 0x41]);
562    }
563
564    #[test]
565    fn a_value_that_overflows_its_type_is_malformed() {
566        assert_eq!(refused::<u16>(&[0x80, 0x80, 0x04]), "65536 overflows u16");
567        assert_eq!(refused::<u32>(&[0x80, 0x80, 0x80, 0x80, 0x10]), "4294967296 overflows u32");
568        assert_eq!(refused::<i16>(&[0x80, 0x80, 0x04]), "32768 overflows i16");
569        assert_eq!(refused::<i32>(&[0x81, 0x80, 0x80, 0x80, 0x10]), "-2147483649 overflows i32");
570        assert_eq!(refused::<bool>(&[2]), "bool 2");
571        assert_eq!(refused::<char>(&[0x80, 0xB0, 0x03]), "0xd800 is not a Unicode scalar value");
572        assert_eq!(refused::<char>(&[0x80, 0x80, 0x44]), "0x110000 is not a Unicode scalar value");
573    }
574
575    #[test]
576    fn integers_are_read_in_their_shortest_form_only() {
577        assert_eq!(refused::<u64>(&[0x80, 0x00]), "an integer is not in its shortest form");
578        assert_eq!(refused::<i64>(&[0x81, 0x00]), "an integer is not in its shortest form");
579        assert_eq!(refused::<u64>(&[0xFF; 10]), "an integer overflows 64 bits");
580        assert_eq!(refused::<u32>(&[0x80]), "the bytes end inside an integer");
581    }
582
583    #[test]
584    fn any_nan_is_written_as_the_canonical_nan() {
585        let odd = f64::from_bits(0xFFF0_0000_0000_0001);
586        assert!(odd.is_nan());
587        assert_eq!(encoded(&odd).0, CANONICAL_NAN.to_le_bytes());
588        assert_eq!(refused::<f64>(&0xFFF0_0000_0000_0001u64.to_le_bytes()), "a NaN other than the canonical one");
589        for v in [0.0, -0.0, 1.5, f64::MIN_POSITIVE, f64::INFINITY, f64::NEG_INFINITY] {
590            let (bytes, strings) = encoded(&v);
591            assert_eq!(decode_all::<f64>("TEST", &bytes, &strings).map(f64::to_bits), Ok(v.to_bits()));
592        }
593        assert_eq!(refused::<f64>(&[0; 7]), "reads past the end: 8 wanted, 7 left");
594    }
595
596    #[test]
597    fn strings_are_indices_into_the_table() {
598        let value = vec!["B".to_owned(), "A".to_owned(), "B".to_owned(), String::new()];
599        let (bytes, strings) = encoded(&value);
600        assert_eq!(bytes, [4, 0, 1, 0, 2]);
601        assert_eq!(strings.iter().collect::<Vec<_>>(), ["B", "A", ""]);
602        round_trip(value);
603        assert_eq!(refused::<String>(&[0]), "string 0 of a table of 0");
604    }
605
606    #[test]
607    fn containers_round_trip() {
608        round_trip(Vec::<u32>::new());
609        round_trip(vec![1u8, 2, 3]);
610        round_trip(vec![Some(-5i64), None]);
611        round_trip([7u16, 300, 65_535]);
612        round_trip(Some(Some(false)));
613        round_trip(Ok::<u8, String>(4));
614        round_trip(Err::<u8, String>("bad".to_owned()));
615        round_trip(Box::new(9u32));
616        round_trip((true, 5u32));
617        round_trip(("X".to_owned(), 1u8, -1i32, None::<u8>));
618        round_trip(BTreeMap::from([(3u32, "c".to_owned()), (1, "a".to_owned())]));
619        assert_eq!(encoded(&vec![1u8, 2]).0, [2, 1, 2]);
620        assert_eq!(encoded(&[1u8, 2]).0, [1, 2]);
621        assert_eq!(encoded(&Some(7u8)).0, [1, 7]);
622        assert_eq!(encoded(&None::<u8>).0, [0]);
623        assert_eq!(encoded(&Err::<u8, u8>(3)).0, [1, 3]);
624        assert_eq!(encoded(&BTreeMap::from([(2u8, 0u8), (1, 9)])).0, [2, 1, 9, 2, 0]);
625    }
626
627    #[test]
628    fn a_bad_container_is_malformed() {
629        assert_eq!(refused::<Vec<u8>>(&[5, 1, 2]), "a count of 5 with 2 bytes left");
630        assert_eq!(refused::<Vec<u8>>(&[0xFF, 0xFF, 0xFF, 0xFF, 0x0F]), "a count of 4294967295 with 0 bytes left");
631        assert_eq!(refused::<Option<u8>>(&[2, 0]), "Option tag 2");
632        assert_eq!(refused::<Result<u8, u8>>(&[2, 0]), "Result tag 2");
633        assert_eq!(refused::<[u8; 3]>(&[1, 2]), "an array of 3 with 2 bytes left");
634        assert_eq!(refused::<BTreeMap<u8, u8>>(&[2, 1, 0, 1, 0]), "map keys that do not strictly ascend");
635        assert_eq!(refused::<BTreeMap<u8, u8>>(&[2, 2, 0, 1, 0]), "map keys that do not strictly ascend");
636        assert_eq!(refused::<u8>(&[1, 2]), "bytes left after the last value: 1");
637        assert_eq!(refused::<u8>(&[]), "reads past the end: 1 wanted, 0 left");
638    }
639
640    #[test]
641    fn an_error_names_the_offset_it_found() {
642        let err = decode_all::<(u8, u8, bool)>("LAYOUT", &[0, 0, 7], &StringTable::default());
643        assert_eq!(err, Err(ModuleError::Malformed { section: "LAYOUT", offset: 2, reason: "bool 7".into() }));
644        assert_eq!(err.unwrap_err().to_string(), "LAYOUT is malformed at byte 2: bool 7");
645    }
646
647    #[test]
648    fn the_same_value_encodes_to_the_same_bytes() {
649        let value = (vec!["Z".to_owned(), "A".to_owned()], BTreeMap::from([(2i32, 'x'), (-1, 'y')]), f64::NAN);
650        assert_eq!(encoded(&value), encoded(&value));
651    }
652}