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 => {
416 let width = match ai {
417 25 => 2,
418 26 => 4,
419 _ => 8,
420 };
421 let bytes = self.take(width)?;
422 let value = match bytes.len() {
423 2 => half_to_f64(u16::from_be_bytes([bytes[0], bytes[1]])),
424 4 => f64::from(f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])),
425 _ => f64::from_be_bytes(bytes.try_into().map_err(|_| DecodeError::Malformed)?),
426 };
427 self.count()?;
428 if value.is_finite() {
429 Ok(Value::Float(value))
430 } else {
431 Err(DecodeError::Malformed)
432 }
433 }
434 0..=24 => {
435 self.argument(ai)?;
436 self.count()?;
437 Err(DecodeError::Malformed)
438 }
439 _ => Err(DecodeError::Malformed),
440 }
441 }
442
443 fn size_hint(&self, count: u64, items_per_element: usize) -> usize {
448 let bytes_left = (self.data.len() - self.pos) / items_per_element;
449 let budget_left = self.budget / items_per_element;
450 count
451 .min(bytes_left as u64)
452 .min(budget_left as u64)
453 .min(MAX_SIZE_HINT as u64) as usize
454 }
455
456 fn list(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
457 if depth >= MAX_NESTING_DEPTH {
458 return Err(DecodeError::NestingTooDeep);
459 }
460 let mut items = Vec::with_capacity(self.size_hint(count, 1));
461 for _ in 0..count {
462 items.push(self.item(depth + 1)?);
463 }
464 Ok(Value::List(items))
465 }
466
467 fn map(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
473 if depth >= MAX_NESTING_DEPTH {
474 return Err(DecodeError::NestingTooDeep);
475 }
476 let hint = self.size_hint(count, 2);
477 let mut pairs = Vec::with_capacity(hint);
478 let mut seen = std::collections::HashSet::with_capacity(hint);
479 for _ in 0..count {
480 let key = self.item(depth + 1)?;
481 let value = self.item(depth + 1)?;
482 let id = match &key {
483 Value::Text(text) => KeyId::Text(text.clone()),
484 Value::Int(n) => KeyId::Int(*n),
485 _ => return Err(DecodeError::BadKey),
486 };
487 if !seen.insert(id) {
488 return Err(DecodeError::DuplicateKey);
489 }
490 pairs.push((key, value));
491 }
492 Ok(Value::Map(pairs))
493 }
494}
495
496fn integer(value: i128, arg: u64) -> Result<Value, DecodeError> {
500 if arg >= 1 << 63 {
501 return Err(DecodeError::IntegerOutOfRange);
502 }
503 Ok(Value::Int(value))
504}
505
506fn half_to_f64(half: u16) -> f64 {
509 let sign = if half >> 15 == 1 { -1.0 } else { 1.0 };
510 let exp = (half >> 10) & 0x1F;
511 let frac = f64::from(half & 0x3FF);
512 match exp {
513 0 => sign * 2f64.powi(-24) * frac,
514 31 if frac == 0.0 => sign * f64::INFINITY,
515 31 => f64::NAN,
516 _ => sign * 2f64.powi(i32::from(exp) - 15) * (1.0 + frac / 1024.0),
517 }
518}
519
520#[cfg(test)]
521mod tests {
522 use super::*;
523
524 fn hex(s: &str) -> Vec<u8> {
527 ::hex::decode(s).expect("valid hex fixture")
528 }
529
530 fn assert_matches_reference(value: Value, expected_hex: &str) {
537 let bytes = encode(&value).expect("encodable fixture");
538 assert_eq!(
539 bytes,
540 hex(expected_hex),
541 "encoding of {value:?} did not match the real macula_cbor_nif output"
542 );
543 let decoded = decode(&bytes).expect("our own output must decode");
547 let re_encoded = encode(&decoded).expect("decoded value must re-encode");
548 assert_eq!(re_encoded, bytes, "encode(decode(bytes)) != bytes");
549 }
550
551 #[test]
552 fn empty_map() {
553 assert_matches_reference(Value::Map(vec![]), "A0");
554 }
555
556 #[test]
557 fn integers_non_negative_minimal_length() {
558 assert_matches_reference(Value::Int(0), "00");
559 assert_matches_reference(Value::Int(23), "17");
560 assert_matches_reference(Value::Int(24), "1818");
561 assert_matches_reference(Value::Int(255), "18FF");
562 assert_matches_reference(Value::Int(256), "190100");
563 assert_matches_reference(Value::Int(65535), "19FFFF");
564 assert_matches_reference(Value::Int(65536), "1A00010000");
565 }
566
567 #[test]
568 fn integers_negative_minimal_length() {
569 assert_matches_reference(Value::Int(-1), "20");
570 assert_matches_reference(Value::Int(-24), "37");
571 assert_matches_reference(Value::Int(-25), "3818");
572 assert_matches_reference(Value::Int(-256), "38FF");
573 }
574
575 #[test]
576 fn integer_out_of_range_is_rejected() {
577 assert_eq!(
579 encode(&Value::Int(u64::MAX as i128 + 1)),
580 Err(IntOutOfRange(u64::MAX as i128 + 1))
581 );
582 let floor = -(1i128 << 64);
584 assert!(encode(&Value::Int(floor)).is_ok());
585 assert!(encode(&Value::Int(floor - 1)).is_err());
586 }
587
588 #[test]
589 fn byte_strings() {
590 assert_matches_reference(Value::Bytes(vec![]), "40");
591 assert_matches_reference(Value::Bytes(b"hello".to_vec()), "4568656C6C6F");
592 }
593
594 #[test]
595 fn text_and_atom_equivalent_encoding() {
596 assert_matches_reference(Value::text("hello"), "6568656C6C6F");
600 assert_matches_reference(Value::text("true"), "6474727565");
601 }
602
603 #[test]
604 fn lists() {
605 assert_matches_reference(Value::List(vec![]), "80");
606 assert_matches_reference(
607 Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
608 "83010203",
609 );
610 }
611
612 #[test]
613 fn floats_always_binary64() {
614 assert_matches_reference(Value::Float(0.0), "FB0000000000000000");
620 assert_matches_reference(Value::Float(12345.6789), "FB40C81CD6E631F8A1");
621 }
622
623 #[test]
624 fn map_keys_sorted_by_encoded_bytes_not_input_order() {
625 assert_matches_reference(
627 Value::Map(vec![
628 (Value::text("b"), Value::Int(2)),
629 (Value::text("a"), Value::Int(1)),
630 ]),
631 "A2616101616202",
632 );
633 }
634
635 #[test]
636 fn map_keys_sorted_lexicographically_same_length() {
637 assert_matches_reference(
638 Value::Map(vec![
639 (Value::text("zebra"), Value::Int(1)),
640 (Value::text("apple"), Value::Int(2)),
641 ]),
642 "A2656170706C6502657A6562726101",
643 );
644 }
645
646 #[test]
647 fn map_keys_shorter_sorts_first_when_prefix() {
648 assert_matches_reference(
652 Value::Map(vec![
653 (Value::text("aa"), Value::Int(1)),
654 (Value::text("a"), Value::Int(2)),
655 (Value::text("aaa"), Value::Int(3)),
656 ]),
657 "A3616102626161016361616103",
658 );
659 }
660
661 #[test]
662 fn null_alone() {
663 assert_matches_reference(Value::Null, "F6");
670 }
671
672 #[test]
673 fn nested_structure_with_null() {
674 assert_matches_reference(
675 Value::Map(vec![
676 (Value::text("name"), Value::text("macula")),
677 (
678 Value::text("nums"),
679 Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
680 ),
681 (Value::text("nil"), Value::Null),
682 ]),
683 "A3636E696CF6646E616D65666D6163756C61646E756D7383010203",
684 );
685 }
686
687 #[test]
688 fn frame_shaped_map() {
689 let node_id: Vec<u8> = (1u8..=32).collect();
690 assert_matches_reference(
691 Value::Map(vec![
692 (Value::text("node_id"), Value::Bytes(node_id)),
693 (Value::text("version"), Value::Int(2)),
694 (Value::text("frame_type"), Value::text("connect")),
695 (Value::text("capabilities"), Value::Int(0)),
696 ]),
697 "A4676E6F64655F696458200102030405060708090A0B0C0D0E0F101112131415161718191A1B1C1D1E1F206776657273696F6E026A6672616D655F7479706567636F6E6E6563746C6361706162696C697469657300",
698 );
699 }
700
701 #[test]
702 fn decode_rejects_trailing_bytes() {
703 assert_eq!(decode(&[0x00, 0xFF]), Err(DecodeError::TrailingBytes));
705 }
706
707 #[test]
717 fn decode_map_with_many_distinct_keys_is_not_quadratic() {
718 let n: i128 = 20_000;
719 let pairs: Vec<(Value, Value)> = (0..n).map(|i| (Value::Int(i), Value::Int(0))).collect();
720 let bytes = encode(&Value::Map(pairs)).expect("encodable");
721
722 let start = std::time::Instant::now();
723 let decoded = decode(&bytes).expect("valid map");
724 let elapsed = start.elapsed();
725
726 match decoded {
727 Value::Map(decoded_pairs) => assert_eq!(decoded_pairs.len(), n as usize),
728 other => panic!("expected a map, got {other:?}"),
729 }
730 assert!(
735 elapsed < std::time::Duration::from_secs(2),
736 "decoding {n} distinct-keyed entries took {elapsed:?} -- \
737 looks like decode_map regressed to O(n^2)"
738 );
739 }
740
741 #[test]
742 fn get_finds_a_field_by_text_key() {
743 let map = Value::Map(vec![(Value::text("a"), Value::Int(1))]);
744 assert_eq!(map.get("a"), Some(&Value::Int(1)));
745 assert_eq!(map.get("missing"), None);
746 }
747
748 #[test]
749 fn get_on_a_non_map_is_none() {
750 assert_eq!(Value::Int(1).get("a"), None);
751 }
752
753 #[test]
754 fn without_removes_only_the_named_keys() {
755 let map = Value::Map(vec![
756 (Value::text("a"), Value::Int(1)),
757 (Value::text("b"), Value::Int(2)),
758 (Value::text("c"), Value::Int(3)),
759 ]);
760 let stripped = map.without(&["b"]);
761 assert_eq!(stripped.get("a"), Some(&Value::Int(1)));
762 assert_eq!(stripped.get("b"), None);
763 assert_eq!(stripped.get("c"), Some(&Value::Int(3)));
764 }
765
766 #[test]
767 fn with_field_replaces_an_existing_key_in_place() {
768 let map =
769 Value::Map(vec![(Value::text("a"), Value::Int(1))]).with_field("a", Value::Int(2));
770 assert_eq!(map.get("a"), Some(&Value::Int(2)));
771 match map {
773 Value::Map(pairs) => assert_eq!(pairs.len(), 1),
774 _ => panic!("expected a map"),
775 }
776 }
777
778 #[test]
779 fn with_field_appends_a_new_key() {
780 let map = Value::Map(vec![]).with_field("a", Value::Int(1));
781 assert_eq!(map.get("a"), Some(&Value::Int(1)));
782 }
783}