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]
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#[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}