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)]
128pub struct IntOutOfRange(pub i128);
129
130impl fmt::Display for IntOutOfRange {
131 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
132 write!(
133 f,
134 "integer {} is outside the encodable range -(2^64)..=u64::MAX",
135 self.0
136 )
137 }
138}
139
140impl std::error::Error for IntOutOfRange {}
141
142#[derive(Debug, Clone, Copy, PartialEq, Eq)]
146pub enum DecodeError {
147 TrailingBytes,
149 BadKey,
151 DuplicateKey,
154 InvalidText,
156 NestingTooDeep,
158 IntegerOutOfRange,
160 TooManyElements,
162 Malformed,
166}
167
168impl fmt::Display for DecodeError {
169 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
170 f.write_str(match self {
171 DecodeError::TrailingBytes => "bytes after the top-level value",
172 DecodeError::BadKey => "a map key that is neither text nor an integer",
173 DecodeError::DuplicateKey => "a duplicate map key",
174 DecodeError::InvalidText => "text that is not valid UTF-8",
175 DecodeError::NestingTooDeep => "arrays and maps nested more than 64 levels",
176 DecodeError::IntegerOutOfRange => "an integer below -2^63 or above 2^63-1",
177 DecodeError::TooManyElements => "more than 131072 items",
178 DecodeError::Malformed => "malformed",
179 })
180 }
181}
182
183impl std::error::Error for DecodeError {}
184
185pub fn encode(value: &Value) -> Result<Vec<u8>, IntOutOfRange> {
189 let mut out = Vec::with_capacity(64);
190 encode_value(value, &mut out)?;
191 Ok(out)
192}
193
194fn encode_value(value: &Value, out: &mut Vec<u8>) -> Result<(), IntOutOfRange> {
195 match value {
196 Value::Int(n) => encode_int(*n, out),
197 Value::Bytes(b) => {
198 encode_head(2, b.len() as u64, out);
199 out.extend_from_slice(b);
200 Ok(())
201 }
202 Value::Text(s) => {
203 let bytes = s.as_bytes();
204 encode_head(3, bytes.len() as u64, out);
205 out.extend_from_slice(bytes);
206 Ok(())
207 }
208 Value::List(items) => {
209 encode_head(4, items.len() as u64, out);
210 for item in items {
211 encode_value(item, out)?;
212 }
213 Ok(())
214 }
215 Value::Map(pairs) => encode_map(pairs, out),
216 Value::Null => {
217 out.push(0xF6); Ok(())
219 }
220 Value::Float(v) => {
221 out.push(0xFB); out.extend_from_slice(&v.to_be_bytes());
223 Ok(())
224 }
225 }
226}
227
228fn encode_int(n: i128, out: &mut Vec<u8>) -> Result<(), IntOutOfRange> {
229 if n >= 0 {
230 if n <= u64::MAX as i128 {
231 encode_head(0, n as u64, out);
232 Ok(())
233 } else {
234 Err(IntOutOfRange(n))
235 }
236 } else {
237 let count = -1i128 - n;
239 if (0..=u64::MAX as i128).contains(&count) {
240 encode_head(1, count as u64, out);
241 Ok(())
242 } else {
243 Err(IntOutOfRange(n))
244 }
245 }
246}
247
248fn encode_map(pairs: &[(Value, Value)], out: &mut Vec<u8>) -> Result<(), IntOutOfRange> {
254 let mut encoded: Vec<(Vec<u8>, Vec<u8>)> = Vec::with_capacity(pairs.len());
255 for (k, v) in pairs {
256 let mut kbuf = Vec::with_capacity(16);
257 encode_value(k, &mut kbuf)?;
258 let mut vbuf = Vec::with_capacity(16);
259 encode_value(v, &mut vbuf)?;
260 encoded.push((kbuf, vbuf));
261 }
262 encoded.sort_by(|a, b| a.0.cmp(&b.0));
263 encode_head(5, encoded.len() as u64, out);
264 for (k, v) in &encoded {
265 out.extend_from_slice(k);
266 out.extend_from_slice(v);
267 }
268 Ok(())
269}
270
271fn encode_head(major: u8, n: u64, out: &mut Vec<u8>) {
272 if n <= 23 {
273 out.push((major << 5) | (n as u8));
274 } else if n <= 0xFF {
275 out.push((major << 5) | 24);
276 out.push(n as u8);
277 } else if n <= 0xFFFF {
278 out.push((major << 5) | 25);
279 out.extend_from_slice(&(n as u16).to_be_bytes());
280 } else if n <= 0xFFFF_FFFF {
281 out.push((major << 5) | 26);
282 out.extend_from_slice(&(n as u32).to_be_bytes());
283 } else {
284 out.push((major << 5) | 27);
285 out.extend_from_slice(&n.to_be_bytes());
286 }
287}
288
289pub const MAX_NESTING_DEPTH: usize = 64;
293
294pub const MAX_ELEMENTS: usize = 131_072;
299
300pub fn decode(bytes: &[u8]) -> Result<Value, DecodeError> {
307 let mut decoder = Decoder {
308 data: bytes,
309 pos: 0,
310 budget: MAX_ELEMENTS,
311 };
312 let value = decoder.item(0)?;
313 if decoder.pos != bytes.len() {
314 return Err(DecodeError::TrailingBytes);
315 }
316 Ok(value)
317}
318
319struct Decoder<'a> {
322 data: &'a [u8],
323 pos: usize,
324 budget: usize,
325}
326
327#[derive(PartialEq, Eq, Hash)]
329enum KeyId {
330 Text(String),
331 Int(i128),
332}
333
334const MAX_SIZE_HINT: usize = 4;
338
339impl Decoder<'_> {
340 fn item(&mut self, depth: usize) -> Result<Value, DecodeError> {
344 let head = self.take(1)?[0];
345 let (major, ai) = (head >> 5, head & 0x1F);
346 if major == 7 {
347 return self.simple_or_float(ai);
348 }
349 let arg = self.argument(ai)?;
350 self.count()?;
351 match major {
352 0 => integer(i128::from(arg), arg),
353 1 => integer(-1 - i128::from(arg), arg),
354 2 => Ok(Value::Bytes(self.take(arg)?.to_vec())),
355 3 => {
356 let bytes = self.take(arg)?;
357 std::str::from_utf8(bytes)
358 .map(|text| Value::Text(text.to_owned()))
359 .map_err(|_| DecodeError::InvalidText)
360 }
361 4 => self.list(arg, depth),
362 5 => self.map(arg, depth),
363 _ => Err(DecodeError::Malformed),
364 }
365 }
366
367 fn count(&mut self) -> Result<(), DecodeError> {
369 if self.budget == 0 {
370 return Err(DecodeError::TooManyElements);
371 }
372 self.budget -= 1;
373 Ok(())
374 }
375
376 fn take(&mut self, n: u64) -> Result<&[u8], DecodeError> {
378 let remaining = (self.data.len() - self.pos) as u64;
379 if n > remaining {
380 return Err(DecodeError::Malformed);
381 }
382 let start = self.pos;
383 self.pos += n as usize;
384 Ok(&self.data[start..self.pos])
385 }
386
387 fn argument(&mut self, ai: u8) -> Result<u64, DecodeError> {
392 let width = match ai {
393 0..=23 => return Ok(u64::from(ai)),
394 24 => 1,
395 25 => 2,
396 26 => 4,
397 27 => 8,
398 _ => return Err(DecodeError::Malformed),
399 };
400 Ok(self
401 .take(width)?
402 .iter()
403 .fold(0u64, |arg, &b| (arg << 8) | u64::from(b)))
404 }
405
406 fn simple_or_float(&mut self, ai: u8) -> Result<Value, DecodeError> {
410 match ai {
411 22 => {
412 self.count()?;
413 Ok(Value::Null)
414 }
415 25..=27 => self.float(ai),
416 0..=24 => {
417 self.argument(ai)?;
418 self.count()?;
419 Err(DecodeError::Malformed)
420 }
421 _ => Err(DecodeError::Malformed),
422 }
423 }
424
425 fn float(&mut self, ai: u8) -> Result<Value, DecodeError> {
429 let width = match ai {
430 25 => 2,
431 26 => 4,
432 _ => 8,
433 };
434 let bytes = self.take(width)?;
435 let value = match bytes.len() {
436 2 => half_to_f64(u16::from_be_bytes([bytes[0], bytes[1]])),
437 4 => f64::from(f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])),
438 _ => f64::from_be_bytes(bytes.try_into().map_err(|_| DecodeError::Malformed)?),
439 };
440 self.count()?;
441 if value.is_finite() {
442 Ok(Value::Float(value))
443 } else {
444 Err(DecodeError::Malformed)
445 }
446 }
447
448 fn size_hint(&self, count: u64, items_per_element: usize) -> usize {
453 let bytes_left = (self.data.len() - self.pos) / items_per_element;
454 let budget_left = self.budget / items_per_element;
455 count
456 .min(bytes_left as u64)
457 .min(budget_left as u64)
458 .min(MAX_SIZE_HINT as u64) as usize
459 }
460
461 fn list(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
462 if depth >= MAX_NESTING_DEPTH {
463 return Err(DecodeError::NestingTooDeep);
464 }
465 let mut items = Vec::with_capacity(self.size_hint(count, 1));
466 for _ in 0..count {
467 items.push(self.item(depth + 1)?);
468 }
469 Ok(Value::List(items))
470 }
471
472 fn map(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
478 if depth >= MAX_NESTING_DEPTH {
479 return Err(DecodeError::NestingTooDeep);
480 }
481 let hint = self.size_hint(count, 2);
482 let mut pairs = Vec::with_capacity(hint);
483 let mut seen = std::collections::HashSet::with_capacity(hint);
484 for _ in 0..count {
485 pairs.push(self.map_entry(depth, &mut seen)?);
486 }
487 Ok(Value::Map(pairs))
488 }
489
490 fn map_entry(
493 &mut self,
494 depth: usize,
495 seen: &mut std::collections::HashSet<KeyId>,
496 ) -> Result<(Value, Value), DecodeError> {
497 let key = self.item(depth + 1)?;
498 let value = self.item(depth + 1)?;
499 if !seen.insert(key_id(&key)?) {
500 return Err(DecodeError::DuplicateKey);
501 }
502 Ok((key, value))
503 }
504}
505
506fn key_id(key: &Value) -> Result<KeyId, DecodeError> {
509 match key {
510 Value::Text(text) => Ok(KeyId::Text(text.clone())),
511 Value::Int(n) => Ok(KeyId::Int(*n)),
512 _ => Err(DecodeError::BadKey),
513 }
514}
515
516fn integer(value: i128, arg: u64) -> Result<Value, DecodeError> {
520 if arg >= 1 << 63 {
521 return Err(DecodeError::IntegerOutOfRange);
522 }
523 Ok(Value::Int(value))
524}
525
526fn half_to_f64(half: u16) -> f64 {
529 let sign = if half >> 15 == 1 { -1.0 } else { 1.0 };
530 let exp = (half >> 10) & 0x1F;
531 let frac = f64::from(half & 0x3FF);
532 match exp {
533 0 => sign * 2f64.powi(-24) * frac,
534 31 if frac == 0.0 => sign * f64::INFINITY,
535 31 => f64::NAN,
536 _ => sign * 2f64.powi(i32::from(exp) - 15) * (1.0 + frac / 1024.0),
537 }
538}
539
540#[cfg(test)]
541mod tests {
542 use super::*;
543
544 fn hex(s: &str) -> Vec<u8> {
547 ::hex::decode(s).expect("valid hex fixture")
548 }
549
550 fn assert_matches_reference(value: Value, expected_hex: &str) {
557 let bytes = encode(&value).expect("encodable fixture");
558 assert_eq!(
559 bytes,
560 hex(expected_hex),
561 "encoding of {value:?} did not match the real macula_cbor_nif output"
562 );
563 let decoded = decode(&bytes).expect("our own output must decode");
567 let re_encoded = encode(&decoded).expect("decoded value must re-encode");
568 assert_eq!(re_encoded, bytes, "encode(decode(bytes)) != bytes");
569 }
570
571 #[test]
572 fn empty_map() {
573 assert_matches_reference(Value::Map(vec![]), "A0");
574 }
575
576 #[test]
577 fn integers_non_negative_minimal_length() {
578 assert_matches_reference(Value::Int(0), "00");
579 assert_matches_reference(Value::Int(23), "17");
580 assert_matches_reference(Value::Int(24), "1818");
581 assert_matches_reference(Value::Int(255), "18FF");
582 assert_matches_reference(Value::Int(256), "190100");
583 assert_matches_reference(Value::Int(65535), "19FFFF");
584 assert_matches_reference(Value::Int(65536), "1A00010000");
585 }
586
587 #[test]
588 fn integers_negative_minimal_length() {
589 assert_matches_reference(Value::Int(-1), "20");
590 assert_matches_reference(Value::Int(-24), "37");
591 assert_matches_reference(Value::Int(-25), "3818");
592 assert_matches_reference(Value::Int(-256), "38FF");
593 }
594
595 #[test]
596 fn integer_out_of_range_is_rejected() {
597 assert_eq!(
599 encode(&Value::Int(u64::MAX as i128 + 1)),
600 Err(IntOutOfRange(u64::MAX as i128 + 1))
601 );
602 let floor = -(1i128 << 64);
604 assert!(encode(&Value::Int(floor)).is_ok());
605 assert!(encode(&Value::Int(floor - 1)).is_err());
606 }
607
608 #[test]
609 fn byte_strings() {
610 assert_matches_reference(Value::Bytes(vec![]), "40");
611 assert_matches_reference(Value::Bytes(b"hello".to_vec()), "4568656C6C6F");
612 }
613
614 #[test]
615 fn text_and_atom_equivalent_encoding() {
616 assert_matches_reference(Value::text("hello"), "6568656C6C6F");
620 assert_matches_reference(Value::text("true"), "6474727565");
621 }
622
623 #[test]
624 fn lists() {
625 assert_matches_reference(Value::List(vec![]), "80");
626 assert_matches_reference(
627 Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
628 "83010203",
629 );
630 }
631
632 #[test]
633 fn floats_always_binary64() {
634 assert_matches_reference(Value::Float(0.0), "FB0000000000000000");
640 assert_matches_reference(Value::Float(12345.6789), "FB40C81CD6E631F8A1");
641 }
642
643 #[test]
644 fn map_keys_sorted_by_encoded_bytes_not_input_order() {
645 assert_matches_reference(
647 Value::Map(vec![
648 (Value::text("b"), Value::Int(2)),
649 (Value::text("a"), Value::Int(1)),
650 ]),
651 "A2616101616202",
652 );
653 }
654
655 #[test]
656 fn map_keys_sorted_lexicographically_same_length() {
657 assert_matches_reference(
658 Value::Map(vec![
659 (Value::text("zebra"), Value::Int(1)),
660 (Value::text("apple"), Value::Int(2)),
661 ]),
662 "A2656170706C6502657A6562726101",
663 );
664 }
665
666 #[test]
667 fn map_keys_shorter_sorts_first_when_prefix() {
668 assert_matches_reference(
672 Value::Map(vec![
673 (Value::text("aa"), Value::Int(1)),
674 (Value::text("a"), Value::Int(2)),
675 (Value::text("aaa"), Value::Int(3)),
676 ]),
677 "A3616102626161016361616103",
678 );
679 }
680
681 #[test]
682 fn null_alone() {
683 assert_matches_reference(Value::Null, "F6");
690 }
691
692 #[test]
693 fn nested_structure_with_null() {
694 assert_matches_reference(
695 Value::Map(vec![
696 (Value::text("name"), Value::text("macula")),
697 (
698 Value::text("nums"),
699 Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
700 ),
701 (Value::text("nil"), Value::Null),
702 ]),
703 "A3636E696CF6646E616D65666D6163756C61646E756D7383010203",
704 );
705 }
706
707 #[test]
708 fn frame_shaped_map() {
709 let node_id: Vec<u8> = (1u8..=32).collect();
710 assert_matches_reference(
711 Value::Map(vec![
712 (Value::text("node_id"), Value::Bytes(node_id)),
713 (Value::text("version"), Value::Int(2)),
714 (Value::text("frame_type"), Value::text("connect")),
715 (Value::text("capabilities"), Value::Int(0)),
716 ]),
717 "A4676E6F64655F696458200102030405060708090A0B0C0D0E0F101112131415161718191A1B1C1D1E1F206776657273696F6E026A6672616D655F7479706567636F6E6E6563746C6361706162696C697469657300",
718 );
719 }
720
721 #[test]
722 fn decode_rejects_trailing_bytes() {
723 assert_eq!(decode(&[0x00, 0xFF]), Err(DecodeError::TrailingBytes));
725 }
726
727 #[test]
737 fn decode_map_with_many_distinct_keys_is_not_quadratic() {
738 let n: i128 = 20_000;
739 let pairs: Vec<(Value, Value)> = (0..n).map(|i| (Value::Int(i), Value::Int(0))).collect();
740 let bytes = encode(&Value::Map(pairs)).expect("encodable");
741
742 let start = std::time::Instant::now();
743 let decoded = decode(&bytes).expect("valid map");
744 let elapsed = start.elapsed();
745
746 match decoded {
747 Value::Map(decoded_pairs) => assert_eq!(decoded_pairs.len(), n as usize),
748 other => panic!("expected a map, got {other:?}"),
749 }
750 assert!(
755 elapsed < std::time::Duration::from_secs(2),
756 "decoding {n} distinct-keyed entries took {elapsed:?} -- \
757 looks like decode_map regressed to O(n^2)"
758 );
759 }
760
761 #[test]
762 fn get_finds_a_field_by_text_key() {
763 let map = Value::Map(vec![(Value::text("a"), Value::Int(1))]);
764 assert_eq!(map.get("a"), Some(&Value::Int(1)));
765 assert_eq!(map.get("missing"), None);
766 }
767
768 #[test]
769 fn get_on_a_non_map_is_none() {
770 assert_eq!(Value::Int(1).get("a"), None);
771 }
772
773 #[test]
774 fn without_removes_only_the_named_keys() {
775 let map = Value::Map(vec![
776 (Value::text("a"), Value::Int(1)),
777 (Value::text("b"), Value::Int(2)),
778 (Value::text("c"), Value::Int(3)),
779 ]);
780 let stripped = map.without(&["b"]);
781 assert_eq!(stripped.get("a"), Some(&Value::Int(1)));
782 assert_eq!(stripped.get("b"), None);
783 assert_eq!(stripped.get("c"), Some(&Value::Int(3)));
784 }
785
786 #[test]
787 fn with_field_replaces_an_existing_key_in_place() {
788 let map =
789 Value::Map(vec![(Value::text("a"), Value::Int(1))]).with_field("a", Value::Int(2));
790 assert_eq!(map.get("a"), Some(&Value::Int(2)));
791 match map {
793 Value::Map(pairs) => assert_eq!(pairs.len(), 1),
794 _ => panic!("expected a map"),
795 }
796 }
797
798 #[test]
799 fn with_field_appends_a_new_key() {
800 let map = Value::Map(vec![]).with_field("a", Value::Int(1));
801 assert_eq!(map.get("a"), Some(&Value::Int(1)));
802 }
803}