1use std::fmt;
54
55#[derive(Debug, Clone, PartialEq)]
58pub enum Value {
59 Int(i128),
62 Bytes(Vec<u8>),
63 Text(String),
66 List(Vec<Value>),
67 Map(Vec<(Value, Value)>),
70 Null,
71 Float(f64),
74}
75
76impl Value {
77 pub fn text(s: impl Into<String>) -> Self {
79 Value::Text(s.into())
80 }
81
82 pub fn get(&self, key: &str) -> Option<&Value> {
86 match self {
87 Value::Map(pairs) => pairs
88 .iter()
89 .find(|(k, _)| matches!(k, Value::Text(t) if t == key))
90 .map(|(_, v)| v),
91 _ => None,
92 }
93 }
94
95 pub fn without(&self, keys: &[&str]) -> Value {
99 match self {
100 Value::Map(pairs) => Value::Map(
101 pairs
102 .iter()
103 .filter(|(k, _)| !matches!(k, Value::Text(t) if keys.contains(&t.as_str())))
104 .cloned()
105 .collect(),
106 ),
107 other => other.clone(),
108 }
109 }
110
111 pub fn with_field(mut self, key: &str, value: Value) -> Value {
114 if let Value::Map(pairs) = &mut self {
115 match pairs
116 .iter_mut()
117 .find(|(k, _)| matches!(k, Value::Text(t) if t == key))
118 {
119 Some(entry) => entry.1 = value,
120 None => pairs.push((Value::text(key), value)),
121 }
122 }
123 self
124 }
125}
126
127#[derive(Debug, Clone, Copy, PartialEq, Eq)]
131pub enum EncodeError {
132 BadKey,
134 DuplicateKey,
136 NestingTooDeep,
138 IntegerOutOfRange,
140 TooManyElements,
142 NonFiniteFloat,
144}
145
146impl fmt::Display for EncodeError {
147 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
148 f.write_str(match self {
149 EncodeError::BadKey => "a map key that is neither text nor an integer",
150 EncodeError::DuplicateKey => "a duplicate map key",
151 EncodeError::NestingTooDeep => "arrays and maps nested more than 64 levels",
152 EncodeError::IntegerOutOfRange => "an integer below -2^63 or above 2^63-1",
153 EncodeError::TooManyElements => "more than 131072 items",
154 EncodeError::NonFiniteFloat => "a float that is NaN or infinite",
155 })
156 }
157}
158
159impl std::error::Error for EncodeError {}
160
161#[derive(Debug, Clone, Copy, PartialEq, Eq)]
165pub enum DecodeError {
166 TrailingBytes,
168 BadKey,
170 DuplicateKey,
173 InvalidText,
175 NestingTooDeep,
177 IntegerOutOfRange,
179 TooManyElements,
181 Malformed,
185}
186
187impl fmt::Display for DecodeError {
188 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
189 f.write_str(match self {
190 DecodeError::TrailingBytes => "bytes after the top-level value",
191 DecodeError::BadKey => "a map key that is neither text nor an integer",
192 DecodeError::DuplicateKey => "a duplicate map key",
193 DecodeError::InvalidText => "text that is not valid UTF-8",
194 DecodeError::NestingTooDeep => "arrays and maps nested more than 64 levels",
195 DecodeError::IntegerOutOfRange => "an integer below -2^63 or above 2^63-1",
196 DecodeError::TooManyElements => "more than 131072 items",
197 DecodeError::Malformed => "malformed",
198 })
199 }
200}
201
202impl std::error::Error for DecodeError {}
203
204pub fn encode(value: &Value) -> Result<Vec<u8>, EncodeError> {
209 let mut encoder = Encoder {
210 out: Vec::with_capacity(64),
211 budget: MAX_ELEMENTS,
212 };
213 encoder.value(value, 0)?;
214 Ok(encoder.out)
215}
216
217struct Encoder {
220 out: Vec<u8>,
221 budget: usize,
222}
223
224impl Encoder {
225 fn value(&mut self, value: &Value, depth: usize) -> Result<(), EncodeError> {
227 if self.budget == 0 {
228 return Err(EncodeError::TooManyElements);
229 }
230 self.budget -= 1;
231 match value {
232 Value::Int(n) => encode_int(*n, &mut self.out),
233 Value::Bytes(b) => {
234 encode_head(2, b.len() as u64, &mut self.out);
235 self.out.extend_from_slice(b);
236 Ok(())
237 }
238 Value::Text(s) => {
239 let bytes = s.as_bytes();
240 encode_head(3, bytes.len() as u64, &mut self.out);
241 self.out.extend_from_slice(bytes);
242 Ok(())
243 }
244 Value::List(items) => self.list(items, depth),
245 Value::Map(pairs) => self.map(pairs, depth),
246 Value::Null => {
247 self.out.push(0xF6); Ok(())
249 }
250 Value::Float(v) if v.is_finite() => {
251 self.out.push(0xFB); self.out.extend_from_slice(&v.to_be_bytes());
253 Ok(())
254 }
255 Value::Float(_) => Err(EncodeError::NonFiniteFloat),
256 }
257 }
258
259 fn list(&mut self, items: &[Value], depth: usize) -> Result<(), EncodeError> {
260 if depth >= MAX_NESTING_DEPTH {
261 return Err(EncodeError::NestingTooDeep);
262 }
263 encode_head(4, items.len() as u64, &mut self.out);
264 for item in items {
265 self.value(item, depth + 1)?;
266 }
267 Ok(())
268 }
269
270 fn map(&mut self, pairs: &[(Value, Value)], depth: usize) -> Result<(), EncodeError> {
278 if depth >= MAX_NESTING_DEPTH {
279 return Err(EncodeError::NestingTooDeep);
280 }
281 let mut encoded: Vec<(Vec<u8>, Vec<u8>)> = Vec::with_capacity(pairs.len());
282 for (k, v) in pairs {
283 encoded.push(self.entry(k, v, depth + 1)?);
284 }
285 encoded.sort_by(|a, b| a.0.cmp(&b.0));
286 if encoded.windows(2).any(|w| w[0].0 == w[1].0) {
287 return Err(EncodeError::DuplicateKey);
288 }
289 encode_head(5, encoded.len() as u64, &mut self.out);
290 for (k, v) in &encoded {
291 self.out.extend_from_slice(k);
292 self.out.extend_from_slice(v);
293 }
294 Ok(())
295 }
296
297 fn entry(
300 &mut self,
301 key: &Value,
302 value: &Value,
303 depth: usize,
304 ) -> Result<(Vec<u8>, Vec<u8>), EncodeError> {
305 let key_bytes = self.piece(key, depth)?;
306 let value_bytes = self.piece(value, depth)?;
307 match key {
308 Value::Text(_) | Value::Int(_) => Ok((key_bytes, value_bytes)),
309 _ => Err(EncodeError::BadKey),
310 }
311 }
312
313 fn piece(&mut self, value: &Value, depth: usize) -> Result<Vec<u8>, EncodeError> {
315 let outer = std::mem::replace(&mut self.out, Vec::with_capacity(16));
316 let written = self.value(value, depth);
317 let piece = std::mem::replace(&mut self.out, outer);
318 written.map(|()| piece)
319 }
320}
321
322fn encode_int(n: i128, out: &mut Vec<u8>) -> Result<(), EncodeError> {
325 if !(i128::from(i64::MIN)..=i128::from(i64::MAX)).contains(&n) {
326 return Err(EncodeError::IntegerOutOfRange);
327 }
328 if n >= 0 {
329 encode_head(0, n as u64, out);
330 } else {
331 encode_head(1, (-1 - n) as u64, out);
332 }
333 Ok(())
334}
335
336fn encode_head(major: u8, n: u64, out: &mut Vec<u8>) {
337 if n <= 23 {
338 out.push((major << 5) | (n as u8));
339 } else if n <= 0xFF {
340 out.push((major << 5) | 24);
341 out.push(n as u8);
342 } else if n <= 0xFFFF {
343 out.push((major << 5) | 25);
344 out.extend_from_slice(&(n as u16).to_be_bytes());
345 } else if n <= 0xFFFF_FFFF {
346 out.push((major << 5) | 26);
347 out.extend_from_slice(&(n as u32).to_be_bytes());
348 } else {
349 out.push((major << 5) | 27);
350 out.extend_from_slice(&n.to_be_bytes());
351 }
352}
353
354pub const MAX_NESTING_DEPTH: usize = 64;
358
359pub const MAX_ELEMENTS: usize = 131_072;
364
365pub fn decode(bytes: &[u8]) -> Result<Value, DecodeError> {
372 let mut decoder = Decoder {
373 data: bytes,
374 pos: 0,
375 budget: MAX_ELEMENTS,
376 };
377 let value = decoder.item(0)?;
378 if decoder.pos != bytes.len() {
379 return Err(DecodeError::TrailingBytes);
380 }
381 Ok(value)
382}
383
384struct Decoder<'a> {
387 data: &'a [u8],
388 pos: usize,
389 budget: usize,
390}
391
392#[derive(PartialEq, Eq, Hash)]
394enum KeyId {
395 Text(String),
396 Int(i128),
397}
398
399const MAX_SIZE_HINT: usize = 4;
403
404impl Decoder<'_> {
405 fn item(&mut self, depth: usize) -> Result<Value, DecodeError> {
409 let head = self.take(1)?[0];
410 let (major, ai) = (head >> 5, head & 0x1F);
411 if major == 7 {
412 return self.simple_or_float(ai);
413 }
414 let arg = self.argument(ai)?;
415 self.count()?;
416 match major {
417 0 => integer(i128::from(arg), arg),
418 1 => integer(-1 - i128::from(arg), arg),
419 2 => Ok(Value::Bytes(self.take(arg)?.to_vec())),
420 3 => {
421 let bytes = self.take(arg)?;
422 std::str::from_utf8(bytes)
423 .map(|text| Value::Text(text.to_owned()))
424 .map_err(|_| DecodeError::InvalidText)
425 }
426 4 => self.list(arg, depth),
427 5 => self.map(arg, depth),
428 _ => Err(DecodeError::Malformed),
429 }
430 }
431
432 fn count(&mut self) -> Result<(), DecodeError> {
434 if self.budget == 0 {
435 return Err(DecodeError::TooManyElements);
436 }
437 self.budget -= 1;
438 Ok(())
439 }
440
441 fn take(&mut self, n: u64) -> Result<&[u8], DecodeError> {
443 let remaining = (self.data.len() - self.pos) as u64;
444 if n > remaining {
445 return Err(DecodeError::Malformed);
446 }
447 let start = self.pos;
448 self.pos += n as usize;
449 Ok(&self.data[start..self.pos])
450 }
451
452 fn argument(&mut self, ai: u8) -> Result<u64, DecodeError> {
457 let width = match ai {
458 0..=23 => return Ok(u64::from(ai)),
459 24 => 1,
460 25 => 2,
461 26 => 4,
462 27 => 8,
463 _ => return Err(DecodeError::Malformed),
464 };
465 Ok(self
466 .take(width)?
467 .iter()
468 .fold(0u64, |arg, &b| (arg << 8) | u64::from(b)))
469 }
470
471 fn simple_or_float(&mut self, ai: u8) -> Result<Value, DecodeError> {
475 match ai {
476 22 => {
477 self.count()?;
478 Ok(Value::Null)
479 }
480 25..=27 => self.float(ai),
481 0..=24 => {
482 self.argument(ai)?;
483 self.count()?;
484 Err(DecodeError::Malformed)
485 }
486 _ => Err(DecodeError::Malformed),
487 }
488 }
489
490 fn float(&mut self, ai: u8) -> Result<Value, DecodeError> {
494 let width = match ai {
495 25 => 2,
496 26 => 4,
497 _ => 8,
498 };
499 let bytes = self.take(width)?;
500 let value = match bytes.len() {
501 2 => half_to_f64(u16::from_be_bytes([bytes[0], bytes[1]])),
502 4 => f64::from(f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])),
503 _ => f64::from_be_bytes(bytes.try_into().map_err(|_| DecodeError::Malformed)?),
504 };
505 self.count()?;
506 if value.is_finite() {
507 Ok(Value::Float(value))
508 } else {
509 Err(DecodeError::Malformed)
510 }
511 }
512
513 fn size_hint(&self, count: u64, items_per_element: usize) -> usize {
518 let bytes_left = (self.data.len() - self.pos) / items_per_element;
519 let budget_left = self.budget / items_per_element;
520 count
521 .min(bytes_left as u64)
522 .min(budget_left as u64)
523 .min(MAX_SIZE_HINT as u64) as usize
524 }
525
526 fn list(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
527 if depth >= MAX_NESTING_DEPTH {
528 return Err(DecodeError::NestingTooDeep);
529 }
530 let mut items = Vec::with_capacity(self.size_hint(count, 1));
531 for _ in 0..count {
532 items.push(self.item(depth + 1)?);
533 }
534 Ok(Value::List(items))
535 }
536
537 fn map(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
543 if depth >= MAX_NESTING_DEPTH {
544 return Err(DecodeError::NestingTooDeep);
545 }
546 let hint = self.size_hint(count, 2);
547 let mut pairs = Vec::with_capacity(hint);
548 let mut seen = std::collections::HashSet::with_capacity(hint);
549 for _ in 0..count {
550 pairs.push(self.map_entry(depth, &mut seen)?);
551 }
552 Ok(Value::Map(pairs))
553 }
554
555 fn map_entry(
558 &mut self,
559 depth: usize,
560 seen: &mut std::collections::HashSet<KeyId>,
561 ) -> Result<(Value, Value), DecodeError> {
562 let key = self.item(depth + 1)?;
563 let value = self.item(depth + 1)?;
564 if !seen.insert(key_id(&key)?) {
565 return Err(DecodeError::DuplicateKey);
566 }
567 Ok((key, value))
568 }
569}
570
571fn key_id(key: &Value) -> Result<KeyId, DecodeError> {
574 match key {
575 Value::Text(text) => Ok(KeyId::Text(text.clone())),
576 Value::Int(n) => Ok(KeyId::Int(*n)),
577 _ => Err(DecodeError::BadKey),
578 }
579}
580
581fn integer(value: i128, arg: u64) -> Result<Value, DecodeError> {
585 if arg >= 1 << 63 {
586 return Err(DecodeError::IntegerOutOfRange);
587 }
588 Ok(Value::Int(value))
589}
590
591fn half_to_f64(half: u16) -> f64 {
594 let sign = if half >> 15 == 1 { -1.0 } else { 1.0 };
595 let exp = (half >> 10) & 0x1F;
596 let frac = f64::from(half & 0x3FF);
597 match exp {
598 0 => sign * 2f64.powi(-24) * frac,
599 31 if frac == 0.0 => sign * f64::INFINITY,
600 31 => f64::NAN,
601 _ => sign * 2f64.powi(i32::from(exp) - 15) * (1.0 + frac / 1024.0),
602 }
603}
604
605#[cfg(test)]
606mod tests {
607 use super::*;
608
609 fn hex(s: &str) -> Vec<u8> {
612 ::hex::decode(s).expect("valid hex fixture")
613 }
614
615 fn assert_matches_reference(value: Value, expected_hex: &str) {
622 let bytes = encode(&value).expect("encodable fixture");
623 assert_eq!(
624 bytes,
625 hex(expected_hex),
626 "encoding of {value:?} did not match the real macula_cbor_nif output"
627 );
628 let decoded = decode(&bytes).expect("our own output must decode");
632 let re_encoded = encode(&decoded).expect("decoded value must re-encode");
633 assert_eq!(re_encoded, bytes, "encode(decode(bytes)) != bytes");
634 }
635
636 #[test]
637 fn empty_map() {
638 assert_matches_reference(Value::Map(vec![]), "A0");
639 }
640
641 #[test]
642 fn integers_non_negative_minimal_length() {
643 assert_matches_reference(Value::Int(0), "00");
644 assert_matches_reference(Value::Int(23), "17");
645 assert_matches_reference(Value::Int(24), "1818");
646 assert_matches_reference(Value::Int(255), "18FF");
647 assert_matches_reference(Value::Int(256), "190100");
648 assert_matches_reference(Value::Int(65535), "19FFFF");
649 assert_matches_reference(Value::Int(65536), "1A00010000");
650 }
651
652 #[test]
653 fn integers_negative_minimal_length() {
654 assert_matches_reference(Value::Int(-1), "20");
655 assert_matches_reference(Value::Int(-24), "37");
656 assert_matches_reference(Value::Int(-25), "3818");
657 assert_matches_reference(Value::Int(-256), "38FF");
658 }
659
660 #[test]
661 fn integer_out_of_range_is_rejected() {
662 let (floor, ceiling) = (i128::from(i64::MIN), i128::from(i64::MAX));
664 assert!(encode(&Value::Int(floor)).is_ok());
665 assert!(encode(&Value::Int(ceiling)).is_ok());
666 assert_eq!(
667 encode(&Value::Int(floor - 1)),
668 Err(EncodeError::IntegerOutOfRange)
669 );
670 assert_eq!(
671 encode(&Value::Int(ceiling + 1)),
672 Err(EncodeError::IntegerOutOfRange)
673 );
674 }
675
676 #[test]
677 fn byte_strings() {
678 assert_matches_reference(Value::Bytes(vec![]), "40");
679 assert_matches_reference(Value::Bytes(b"hello".to_vec()), "4568656C6C6F");
680 }
681
682 #[test]
683 fn text_and_atom_equivalent_encoding() {
684 assert_matches_reference(Value::text("hello"), "6568656C6C6F");
688 assert_matches_reference(Value::text("true"), "6474727565");
689 }
690
691 #[test]
692 fn lists() {
693 assert_matches_reference(Value::List(vec![]), "80");
694 assert_matches_reference(
695 Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
696 "83010203",
697 );
698 }
699
700 #[test]
701 fn floats_always_binary64() {
702 assert_matches_reference(Value::Float(0.0), "FB0000000000000000");
708 assert_matches_reference(Value::Float(12345.6789), "FB40C81CD6E631F8A1");
709 }
710
711 #[test]
712 fn map_keys_sorted_by_encoded_bytes_not_input_order() {
713 assert_matches_reference(
715 Value::Map(vec![
716 (Value::text("b"), Value::Int(2)),
717 (Value::text("a"), Value::Int(1)),
718 ]),
719 "A2616101616202",
720 );
721 }
722
723 #[test]
724 fn map_keys_sorted_lexicographically_same_length() {
725 assert_matches_reference(
726 Value::Map(vec![
727 (Value::text("zebra"), Value::Int(1)),
728 (Value::text("apple"), Value::Int(2)),
729 ]),
730 "A2656170706C6502657A6562726101",
731 );
732 }
733
734 #[test]
735 fn map_keys_shorter_sorts_first_when_prefix() {
736 assert_matches_reference(
740 Value::Map(vec![
741 (Value::text("aa"), Value::Int(1)),
742 (Value::text("a"), Value::Int(2)),
743 (Value::text("aaa"), Value::Int(3)),
744 ]),
745 "A3616102626161016361616103",
746 );
747 }
748
749 #[test]
750 fn null_alone() {
751 assert_matches_reference(Value::Null, "F6");
758 }
759
760 #[test]
761 fn nested_structure_with_null() {
762 assert_matches_reference(
763 Value::Map(vec![
764 (Value::text("name"), Value::text("macula")),
765 (
766 Value::text("nums"),
767 Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
768 ),
769 (Value::text("nil"), Value::Null),
770 ]),
771 "A3636E696CF6646E616D65666D6163756C61646E756D7383010203",
772 );
773 }
774
775 #[test]
776 fn frame_shaped_map() {
777 let node_id: Vec<u8> = (1u8..=32).collect();
778 assert_matches_reference(
779 Value::Map(vec![
780 (Value::text("node_id"), Value::Bytes(node_id)),
781 (Value::text("version"), Value::Int(2)),
782 (Value::text("frame_type"), Value::text("connect")),
783 (Value::text("capabilities"), Value::Int(0)),
784 ]),
785 "A4676E6F64655F696458200102030405060708090A0B0C0D0E0F101112131415161718191A1B1C1D1E1F206776657273696F6E026A6672616D655F7479706567636F6E6E6563746C6361706162696C697469657300",
786 );
787 }
788
789 #[test]
790 fn decode_rejects_trailing_bytes() {
791 assert_eq!(decode(&[0x00, 0xFF]), Err(DecodeError::TrailingBytes));
793 }
794
795 #[test]
805 fn decode_map_with_many_distinct_keys_is_not_quadratic() {
806 let n: i128 = 20_000;
807 let pairs: Vec<(Value, Value)> = (0..n).map(|i| (Value::Int(i), Value::Int(0))).collect();
808 let bytes = encode(&Value::Map(pairs)).expect("encodable");
809
810 let start = std::time::Instant::now();
811 let decoded = decode(&bytes).expect("valid map");
812 let elapsed = start.elapsed();
813
814 match decoded {
815 Value::Map(decoded_pairs) => assert_eq!(decoded_pairs.len(), n as usize),
816 other => panic!("expected a map, got {other:?}"),
817 }
818 assert!(
823 elapsed < std::time::Duration::from_secs(2),
824 "decoding {n} distinct-keyed entries took {elapsed:?} -- \
825 looks like decode_map regressed to O(n^2)"
826 );
827 }
828
829 #[test]
830 fn get_finds_a_field_by_text_key() {
831 let map = Value::Map(vec![(Value::text("a"), Value::Int(1))]);
832 assert_eq!(map.get("a"), Some(&Value::Int(1)));
833 assert_eq!(map.get("missing"), None);
834 }
835
836 #[test]
837 fn get_on_a_non_map_is_none() {
838 assert_eq!(Value::Int(1).get("a"), None);
839 }
840
841 #[test]
842 fn without_removes_only_the_named_keys() {
843 let map = Value::Map(vec![
844 (Value::text("a"), Value::Int(1)),
845 (Value::text("b"), Value::Int(2)),
846 (Value::text("c"), Value::Int(3)),
847 ]);
848 let stripped = map.without(&["b"]);
849 assert_eq!(stripped.get("a"), Some(&Value::Int(1)));
850 assert_eq!(stripped.get("b"), None);
851 assert_eq!(stripped.get("c"), Some(&Value::Int(3)));
852 }
853
854 #[test]
855 fn with_field_replaces_an_existing_key_in_place() {
856 let map =
857 Value::Map(vec![(Value::text("a"), Value::Int(1))]).with_field("a", Value::Int(2));
858 assert_eq!(map.get("a"), Some(&Value::Int(2)));
859 match map {
861 Value::Map(pairs) => assert_eq!(pairs.len(), 1),
862 _ => panic!("expected a map"),
863 }
864 }
865
866 #[test]
867 fn with_field_appends_a_new_key() {
868 let map = Value::Map(vec![]).with_field("a", Value::Int(1));
869 assert_eq!(map.get("a"), Some(&Value::Int(1)));
870 }
871}