1use 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#[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 pub fn string(&mut self, text: &str) {
51 let index = self.strings.intern(text);
52 self.count(index);
53 }
54
55 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#[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 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 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
153pub 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
161pub 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#[macro_export]
438macro_rules! codec_struct {
439 ($ty:ident { $($field:ident),* $(,)? }) => {
440 $crate::codec_struct!($ty { $($field),* } check $crate::module::codec::unchecked);
441 };
442 ($ty:ident { $($field:ident),* $(,)? } check $check:path) => {
443 $crate::codec_struct!($ty { $($field),* } default {} check $check);
444 };
445 ($ty:ident { $($field:ident),* $(,)? } default { $($left:ident),* $(,)? } check $check:path) => {
446 impl $crate::module::codec::Encode for $ty {
447 fn encode(&self, w: &mut $crate::module::codec::Writer) {
448 let $ty { $($field,)* $($left: _,)* } = self;
449 $($crate::module::codec::Encode::encode($field, w);)*
450 }
451 }
452
453 impl $crate::module::codec::Decode for $ty {
454 fn decode(
455 r: &mut $crate::module::codec::Reader<'_>,
456 ) -> ::core::result::Result<Self, $crate::module::ModuleError> {
457 let at = r.position();
458 $(let $field = $crate::module::codec::Decode::decode(r)?;)*
459 let value = $ty { $($field,)* $($left: ::core::default::Default::default(),)* };
460 $check(&value).map_err(|reason| r.malformed(at, reason))?;
461 ::core::result::Result::Ok(value)
462 }
463 }
464 };
465}
466
467#[macro_export]
469macro_rules! codec_enum {
470 ($ty:ident {
471 $($variant:ident $({ $($field:ident),* $(,)? })? $(( $($elem:ident),* $(,)? ))? = $tag:literal),* $(,)?
472 }) => {
473 impl $crate::module::codec::Encode for $ty {
474 fn encode(&self, w: &mut $crate::module::codec::Writer) {
475 match self {
476 $($ty::$variant $({ $($field),* })? $(( $($elem),* ))? => {
477 w.leb($tag);
478 $($($crate::module::codec::Encode::encode($field, w);)*)?
479 $($($crate::module::codec::Encode::encode($elem, w);)*)?
480 })*
481 }
482 }
483 }
484
485 impl $crate::module::codec::Decode for $ty {
486 fn decode(
487 r: &mut $crate::module::codec::Reader<'_>,
488 ) -> ::core::result::Result<Self, $crate::module::ModuleError> {
489 let at = r.position();
490 match r.leb()? {
491 $($tag => {
492 $($(let $field = $crate::module::codec::Decode::decode(r)?;)*)?
493 $($(let $elem = $crate::module::codec::Decode::decode(r)?;)*)?
494 ::core::result::Result::Ok($ty::$variant $({ $($field),* })? $(( $($elem),* ))?)
495 })*
496 tag => ::core::result::Result::Err(r.malformed(at, format!("{} has no tag {tag}", stringify!($ty)))),
497 }
498 }
499 }
500 };
501}
502
503#[cfg(test)]
504mod tests {
505 use super::*;
506
507 fn encoded<T: Encode>(value: &T) -> (Vec<u8>, StringTable) {
508 let mut w = Writer::new();
509 value.encode(&mut w);
510 (w.take(), w.strings().clone())
511 }
512
513 fn round_trip<T: Encode + Decode + PartialEq + std::fmt::Debug>(value: T) {
514 let (bytes, strings) = encoded(&value);
515 assert_eq!(decode_all::<T>("TEST", &bytes, &strings), Ok(value));
516 }
517
518 fn refused<T: Decode + std::fmt::Debug>(bytes: &[u8]) -> String {
519 match decode_all::<T>("TEST", bytes, &StringTable::default()) {
520 Err(ModuleError::Malformed { section: "TEST", reason, .. }) => reason,
521 other => panic!("{bytes:02X?} decoded as {other:?}"),
522 }
523 }
524
525 #[test]
526 fn integers_round_trip_at_their_limits() {
527 for v in [0, 1, 0x7F, 0x80, u8::MAX] {
528 round_trip(v);
529 }
530 for v in [0, 0x80, u16::MAX] {
531 round_trip(v);
532 }
533 for v in [0, 0x3FFF, 0x4000, u32::MAX] {
534 round_trip(v);
535 }
536 for v in [0, u64::from(u32::MAX) + 1, u64::MAX] {
537 round_trip(v);
538 }
539 for v in [0, usize::MAX] {
540 round_trip(v);
541 }
542 for v in [i16::MIN, -1, 0, 1, i16::MAX] {
543 round_trip(v);
544 }
545 for v in [i32::MIN, -64, 63, i32::MAX] {
546 round_trip(v);
547 }
548 for v in [i64::MIN, -1, 0, i64::MAX] {
549 round_trip(v);
550 }
551 round_trip(false);
552 round_trip(true);
553 for c in ['\0', 'A', 'é', '€', '\u{10FFFF}'] {
554 round_trip(c);
555 }
556 }
557
558 #[test]
559 fn integers_have_the_documented_bytes() {
560 assert_eq!(encoded(&1140u16).0, [0xF4, 0x08]);
561 assert_eq!(encoded(&-1i32).0, [0x01]);
562 assert_eq!(encoded(&-65i64).0, [0x81, 0x01]);
563 assert_eq!(encoded(&200u8).0, [200]);
564 assert_eq!(encoded(&true).0, [1]);
565 assert_eq!(encoded(&'€').0, [0xAC, 0x41]);
566 }
567
568 #[test]
569 fn a_value_that_overflows_its_type_is_malformed() {
570 assert_eq!(refused::<u16>(&[0x80, 0x80, 0x04]), "65536 overflows u16");
571 assert_eq!(refused::<u32>(&[0x80, 0x80, 0x80, 0x80, 0x10]), "4294967296 overflows u32");
572 assert_eq!(refused::<i16>(&[0x80, 0x80, 0x04]), "32768 overflows i16");
573 assert_eq!(refused::<i32>(&[0x81, 0x80, 0x80, 0x80, 0x10]), "-2147483649 overflows i32");
574 assert_eq!(refused::<bool>(&[2]), "bool 2");
575 assert_eq!(refused::<char>(&[0x80, 0xB0, 0x03]), "0xd800 is not a Unicode scalar value");
576 assert_eq!(refused::<char>(&[0x80, 0x80, 0x44]), "0x110000 is not a Unicode scalar value");
577 }
578
579 #[test]
580 fn integers_are_read_in_their_shortest_form_only() {
581 assert_eq!(refused::<u64>(&[0x80, 0x00]), "an integer is not in its shortest form");
582 assert_eq!(refused::<i64>(&[0x81, 0x00]), "an integer is not in its shortest form");
583 assert_eq!(refused::<u64>(&[0xFF; 10]), "an integer overflows 64 bits");
584 assert_eq!(refused::<u32>(&[0x80]), "the bytes end inside an integer");
585 }
586
587 #[test]
588 fn any_nan_is_written_as_the_canonical_nan() {
589 let odd = f64::from_bits(0xFFF0_0000_0000_0001);
590 assert!(odd.is_nan());
591 assert_eq!(encoded(&odd).0, CANONICAL_NAN.to_le_bytes());
592 assert_eq!(refused::<f64>(&0xFFF0_0000_0000_0001u64.to_le_bytes()), "a NaN other than the canonical one");
593 for v in [0.0, -0.0, 1.5, f64::MIN_POSITIVE, f64::INFINITY, f64::NEG_INFINITY] {
594 let (bytes, strings) = encoded(&v);
595 assert_eq!(decode_all::<f64>("TEST", &bytes, &strings).map(f64::to_bits), Ok(v.to_bits()));
596 }
597 assert_eq!(refused::<f64>(&[0; 7]), "reads past the end: 8 wanted, 7 left");
598 }
599
600 #[test]
601 fn strings_are_indices_into_the_table() {
602 let value = vec!["B".to_owned(), "A".to_owned(), "B".to_owned(), String::new()];
603 let (bytes, strings) = encoded(&value);
604 assert_eq!(bytes, [4, 0, 1, 0, 2]);
605 assert_eq!(strings.iter().collect::<Vec<_>>(), ["B", "A", ""]);
606 round_trip(value);
607 assert_eq!(refused::<String>(&[0]), "string 0 of a table of 0");
608 }
609
610 #[test]
611 fn containers_round_trip() {
612 round_trip(Vec::<u32>::new());
613 round_trip(vec![1u8, 2, 3]);
614 round_trip(vec![Some(-5i64), None]);
615 round_trip([7u16, 300, 65_535]);
616 round_trip(Some(Some(false)));
617 round_trip(Ok::<u8, String>(4));
618 round_trip(Err::<u8, String>("bad".to_owned()));
619 round_trip(Box::new(9u32));
620 round_trip((true, 5u32));
621 round_trip(("X".to_owned(), 1u8, -1i32, None::<u8>));
622 round_trip(BTreeMap::from([(3u32, "c".to_owned()), (1, "a".to_owned())]));
623 assert_eq!(encoded(&vec![1u8, 2]).0, [2, 1, 2]);
624 assert_eq!(encoded(&[1u8, 2]).0, [1, 2]);
625 assert_eq!(encoded(&Some(7u8)).0, [1, 7]);
626 assert_eq!(encoded(&None::<u8>).0, [0]);
627 assert_eq!(encoded(&Err::<u8, u8>(3)).0, [1, 3]);
628 assert_eq!(encoded(&BTreeMap::from([(2u8, 0u8), (1, 9)])).0, [2, 1, 9, 2, 0]);
629 }
630
631 #[test]
632 fn a_bad_container_is_malformed() {
633 assert_eq!(refused::<Vec<u8>>(&[5, 1, 2]), "a count of 5 with 2 bytes left");
634 assert_eq!(refused::<Vec<u8>>(&[0xFF, 0xFF, 0xFF, 0xFF, 0x0F]), "a count of 4294967295 with 0 bytes left");
635 assert_eq!(refused::<Option<u8>>(&[2, 0]), "Option tag 2");
636 assert_eq!(refused::<Result<u8, u8>>(&[2, 0]), "Result tag 2");
637 assert_eq!(refused::<[u8; 3]>(&[1, 2]), "an array of 3 with 2 bytes left");
638 assert_eq!(refused::<BTreeMap<u8, u8>>(&[2, 1, 0, 1, 0]), "map keys that do not strictly ascend");
639 assert_eq!(refused::<BTreeMap<u8, u8>>(&[2, 2, 0, 1, 0]), "map keys that do not strictly ascend");
640 assert_eq!(refused::<u8>(&[1, 2]), "bytes left after the last value: 1");
641 assert_eq!(refused::<u8>(&[]), "reads past the end: 1 wanted, 0 left");
642 }
643
644 #[test]
645 fn an_error_names_the_offset_it_found() {
646 let err = decode_all::<(u8, u8, bool)>("LAYOUT", &[0, 0, 7], &StringTable::default());
647 assert_eq!(err, Err(ModuleError::Malformed { section: "LAYOUT", offset: 2, reason: "bool 7".into() }));
648 assert_eq!(err.unwrap_err().to_string(), "LAYOUT is malformed at byte 2: bool 7");
649 }
650
651 #[test]
652 fn the_same_value_encodes_to_the_same_bytes() {
653 let value = (vec!["Z".to_owned(), "A".to_owned()], BTreeMap::from([(2i32, 'x'), (-1, 'y')]), f64::NAN);
654 assert_eq!(encoded(&value), encoded(&value));
655 }
656}