use std::fmt;
#[derive(Debug, Clone, PartialEq)]
pub enum Value {
Int(i128),
Bytes(Vec<u8>),
Text(String),
List(Vec<Value>),
Map(Vec<(Value, Value)>),
Null,
Float(f64),
}
impl Value {
pub fn text(s: impl Into<String>) -> Self {
Value::Text(s.into())
}
pub fn get(&self, key: &str) -> Option<&Value> {
match self {
Value::Map(pairs) => pairs
.iter()
.find(|(k, _)| matches!(k, Value::Text(t) if t == key))
.map(|(_, v)| v),
_ => None,
}
}
pub fn without(&self, keys: &[&str]) -> Value {
match self {
Value::Map(pairs) => Value::Map(
pairs
.iter()
.filter(|(k, _)| !matches!(k, Value::Text(t) if keys.contains(&t.as_str())))
.cloned()
.collect(),
),
other => other.clone(),
}
}
pub fn with_field(mut self, key: &str, value: Value) -> Value {
if let Value::Map(pairs) = &mut self {
match pairs
.iter_mut()
.find(|(k, _)| matches!(k, Value::Text(t) if t == key))
{
Some(entry) => entry.1 = value,
None => pairs.push((Value::text(key), value)),
}
}
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct IntOutOfRange(pub i128);
impl fmt::Display for IntOutOfRange {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"integer {} is outside the encodable range -(2^64)..=u64::MAX",
self.0
)
}
}
impl std::error::Error for IntOutOfRange {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DecodeError {
Truncated,
UnsupportedMajorType(u8),
UnsupportedAdditionalInfo(u8),
UnsupportedAdditionalInfoEncoding(u8),
TrailingBytes,
InvalidUtf8,
UnrepresentableFloat,
NestingTooDeep,
}
impl fmt::Display for DecodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
DecodeError::Truncated => write!(f, "truncated input"),
DecodeError::UnsupportedMajorType(m) => {
write!(
f,
"unsupported major type {m} (only 0-5 and 7 are valid here)"
)
}
DecodeError::UnsupportedAdditionalInfo(ai) => {
write!(f, "unsupported major-7 additional info {ai}")
}
DecodeError::UnsupportedAdditionalInfoEncoding(ai) => {
write!(
f,
"unsupported additional-info encoding {ai} (28-31 are reserved)"
)
}
DecodeError::TrailingBytes => write!(f, "trailing bytes after the top-level value"),
DecodeError::InvalidUtf8 => write!(f, "text value was not valid UTF-8"),
DecodeError::UnrepresentableFloat => {
write!(f, "half-float NaN/infinity has no f64 representation here")
}
DecodeError::NestingTooDeep => {
write!(f, "list/map nesting exceeds {MAX_NESTING_DEPTH} levels")
}
}
}
}
impl std::error::Error for DecodeError {}
pub fn encode(value: &Value) -> Result<Vec<u8>, IntOutOfRange> {
let mut out = Vec::with_capacity(64);
encode_value(value, &mut out)?;
Ok(out)
}
fn encode_value(value: &Value, out: &mut Vec<u8>) -> Result<(), IntOutOfRange> {
match value {
Value::Int(n) => encode_int(*n, out),
Value::Bytes(b) => {
encode_head(2, b.len() as u64, out);
out.extend_from_slice(b);
Ok(())
}
Value::Text(s) => {
let bytes = s.as_bytes();
encode_head(3, bytes.len() as u64, out);
out.extend_from_slice(bytes);
Ok(())
}
Value::List(items) => {
encode_head(4, items.len() as u64, out);
for item in items {
encode_value(item, out)?;
}
Ok(())
}
Value::Map(pairs) => encode_map(pairs, out),
Value::Null => {
out.push(0xF6); Ok(())
}
Value::Float(v) => {
out.push(0xFB); out.extend_from_slice(&v.to_be_bytes());
Ok(())
}
}
}
fn encode_int(n: i128, out: &mut Vec<u8>) -> Result<(), IntOutOfRange> {
if n >= 0 {
if n <= u64::MAX as i128 {
encode_head(0, n as u64, out);
Ok(())
} else {
Err(IntOutOfRange(n))
}
} else {
let count = -1i128 - n;
if (0..=u64::MAX as i128).contains(&count) {
encode_head(1, count as u64, out);
Ok(())
} else {
Err(IntOutOfRange(n))
}
}
}
fn encode_map(pairs: &[(Value, Value)], out: &mut Vec<u8>) -> Result<(), IntOutOfRange> {
let mut encoded: Vec<(Vec<u8>, Vec<u8>)> = Vec::with_capacity(pairs.len());
for (k, v) in pairs {
let mut kbuf = Vec::with_capacity(16);
encode_value(k, &mut kbuf)?;
let mut vbuf = Vec::with_capacity(16);
encode_value(v, &mut vbuf)?;
encoded.push((kbuf, vbuf));
}
encoded.sort_by(|a, b| a.0.cmp(&b.0));
encode_head(5, encoded.len() as u64, out);
for (k, v) in &encoded {
out.extend_from_slice(k);
out.extend_from_slice(v);
}
Ok(())
}
fn encode_head(major: u8, n: u64, out: &mut Vec<u8>) {
if n <= 23 {
out.push((major << 5) | (n as u8));
} else if n <= 0xFF {
out.push((major << 5) | 24);
out.push(n as u8);
} else if n <= 0xFFFF {
out.push((major << 5) | 25);
out.extend_from_slice(&(n as u16).to_be_bytes());
} else if n <= 0xFFFF_FFFF {
out.push((major << 5) | 26);
out.extend_from_slice(&(n as u32).to_be_bytes());
} else {
out.push((major << 5) | 27);
out.extend_from_slice(&n.to_be_bytes());
}
}
pub const MAX_NESTING_DEPTH: usize = 128;
pub fn decode(bytes: &[u8]) -> Result<Value, DecodeError> {
let (value, _canonical_bytes, pos) = decode_one(bytes, 0, 0, false)?;
if pos != bytes.len() {
return Err(DecodeError::TrailingBytes);
}
Ok(value)
}
fn need(buf: &[u8], pos: usize, n: usize) -> Result<(), DecodeError> {
match pos.checked_add(n) {
Some(end) if end <= buf.len() => Ok(()),
_ => Err(DecodeError::Truncated),
}
}
fn decode_one(
buf: &[u8],
pos: usize,
depth: usize,
need_canon: bool,
) -> Result<(Value, Vec<u8>, usize), DecodeError> {
if depth > MAX_NESTING_DEPTH {
return Err(DecodeError::NestingTooDeep);
}
need(buf, pos, 1)?;
let byte0 = buf[pos];
let major = byte0 >> 5;
let ai = byte0 & 0x1F;
if major == 7 {
let (value, next) = decode_major7(buf, pos, ai)?;
return Ok(scalar_canonical_bytes(value, next, need_canon));
}
let (n, next) = decode_count(buf, pos + 1, ai)?;
match major {
0 => Ok(scalar_canonical_bytes(
Value::Int(n as i128),
next,
need_canon,
)),
1 => Ok(scalar_canonical_bytes(
Value::Int(-1i128 - n as i128),
next,
need_canon,
)),
2 => {
let len = n as usize;
need(buf, next, len)?;
let value = Value::Bytes(buf[next..next + len].to_vec());
Ok(scalar_canonical_bytes(value, next + len, need_canon))
}
3 => {
let len = n as usize;
need(buf, next, len)?;
let text = String::from_utf8(buf[next..next + len].to_vec())
.map_err(|_| DecodeError::InvalidUtf8)?;
Ok(scalar_canonical_bytes(
Value::Text(text),
next + len,
need_canon,
))
}
4 => decode_list(buf, next, n, depth + 1, need_canon),
5 => decode_map(buf, next, n, depth + 1, need_canon),
_ => Err(DecodeError::UnsupportedMajorType(major)),
}
}
fn scalar_canonical_bytes(value: Value, next: usize, need_canon: bool) -> (Value, Vec<u8>, usize) {
if need_canon {
with_canonical_bytes(value, next)
} else {
(value, Vec::new(), next)
}
}
fn with_canonical_bytes(value: Value, next: usize) -> (Value, Vec<u8>, usize) {
let mut canon = Vec::new();
encode_value(&value, &mut canon).expect("a value produced by this decoder is always encodable");
(value, canon, next)
}
fn decode_count(buf: &[u8], pos: usize, ai: u8) -> Result<(u64, usize), DecodeError> {
match ai {
0..=23 => Ok((ai as u64, pos)),
24 => {
need(buf, pos, 1)?;
Ok((buf[pos] as u64, pos + 1))
}
25 => {
need(buf, pos, 2)?;
Ok((u16::from_be_bytes([buf[pos], buf[pos + 1]]) as u64, pos + 2))
}
26 => {
need(buf, pos, 4)?;
let b: [u8; 4] = buf[pos..pos + 4].try_into().expect("checked len");
Ok((u32::from_be_bytes(b) as u64, pos + 4))
}
27 => {
need(buf, pos, 8)?;
let b: [u8; 8] = buf[pos..pos + 8].try_into().expect("checked len");
Ok((u64::from_be_bytes(b), pos + 8))
}
28..=31 => Err(DecodeError::UnsupportedAdditionalInfoEncoding(ai)),
_ => unreachable!("additional info is a 5-bit field, 0..=31"),
}
}
fn decode_major7(buf: &[u8], pos: usize, ai: u8) -> Result<(Value, usize), DecodeError> {
match ai {
22 => Ok((Value::Null, pos + 1)),
25 => {
need(buf, pos + 1, 2)?;
let half = u16::from_be_bytes([buf[pos + 1], buf[pos + 2]]);
Ok((Value::Float(half_to_f64(half)?), pos + 3))
}
26 => {
need(buf, pos + 1, 4)?;
let b: [u8; 4] = buf[pos + 1..pos + 5].try_into().expect("checked len");
Ok((Value::Float(f32::from_be_bytes(b) as f64), pos + 5))
}
27 => {
need(buf, pos + 1, 8)?;
let b: [u8; 8] = buf[pos + 1..pos + 9].try_into().expect("checked len");
Ok((Value::Float(f64::from_be_bytes(b)), pos + 9))
}
_ => Err(DecodeError::UnsupportedAdditionalInfo(ai)),
}
}
fn decode_list(
buf: &[u8],
mut pos: usize,
count: u64,
depth: usize,
need_canon: bool,
) -> Result<(Value, Vec<u8>, usize), DecodeError> {
let mut items = Vec::with_capacity(count.min(1024) as usize);
let mut canon = Vec::new();
if need_canon {
encode_head(4, count, &mut canon);
}
for _ in 0..count {
let (item, item_canon, next) = decode_one(buf, pos, depth, need_canon)?;
if need_canon {
canon.extend_from_slice(&item_canon);
}
items.push(item);
pos = next;
}
Ok((Value::List(items), canon, pos))
}
fn decode_map(
buf: &[u8],
mut pos: usize,
count: u64,
depth: usize,
need_canon: bool,
) -> Result<(Value, Vec<u8>, usize), DecodeError> {
let capacity = count.min(1024) as usize;
let mut pairs: Vec<(Value, Value)> = Vec::with_capacity(capacity);
let mut index_of_key: std::collections::HashMap<Vec<u8>, usize> =
std::collections::HashMap::with_capacity(capacity);
let mut vals_canon: Vec<Vec<u8>> = Vec::with_capacity(if need_canon { capacity } else { 0 });
for _ in 0..count {
let (k, key_canon, next1) = decode_one(buf, pos, depth, true)?;
let (v, val_canon, next2) = decode_one(buf, next1, depth, need_canon)?;
pos = next2;
use std::collections::hash_map::Entry;
match index_of_key.entry(key_canon) {
Entry::Occupied(e) => {
let i = *e.get();
pairs[i].1 = v;
if need_canon {
vals_canon[i] = val_canon;
}
}
Entry::Vacant(e) => {
e.insert(pairs.len());
pairs.push((k, v));
if need_canon {
vals_canon.push(val_canon);
}
}
}
}
if !need_canon {
return Ok((Value::Map(pairs), Vec::new(), pos));
}
let mut order: Vec<(&Vec<u8>, usize)> = index_of_key.iter().map(|(k, &i)| (k, i)).collect();
order.sort_by(|a, b| a.0.cmp(b.0));
let mut canon = Vec::new();
encode_head(5, order.len() as u64, &mut canon);
for (k, i) in order {
canon.extend_from_slice(k);
canon.extend_from_slice(&vals_canon[i]);
}
Ok((Value::Map(pairs), canon, pos))
}
fn half_to_f64(half: u16) -> Result<f64, DecodeError> {
let sign: f64 = if (half >> 15) & 1 == 1 { -1.0 } else { 1.0 };
let exp = (half >> 10) & 0x1F;
let frac = (half & 0x3FF) as f64;
match exp {
0 => Ok(sign * 2f64.powi(-14) * (frac / 1024.0)),
1..=30 => Ok(sign * 2f64.powi(exp as i32 - 15) * (1.0 + frac / 1024.0)),
_ => Err(DecodeError::UnrepresentableFloat),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn hex(s: &str) -> Vec<u8> {
::hex::decode(s).expect("valid hex fixture")
}
fn assert_matches_reference(value: Value, expected_hex: &str) {
let bytes = encode(&value).expect("encodable fixture");
assert_eq!(
bytes,
hex(expected_hex),
"encoding of {value:?} did not match the real macula_cbor_nif output"
);
let decoded = decode(&bytes).expect("our own output must decode");
let re_encoded = encode(&decoded).expect("decoded value must re-encode");
assert_eq!(re_encoded, bytes, "encode(decode(bytes)) != bytes");
}
#[test]
fn empty_map() {
assert_matches_reference(Value::Map(vec![]), "A0");
}
#[test]
fn integers_non_negative_minimal_length() {
assert_matches_reference(Value::Int(0), "00");
assert_matches_reference(Value::Int(23), "17");
assert_matches_reference(Value::Int(24), "1818");
assert_matches_reference(Value::Int(255), "18FF");
assert_matches_reference(Value::Int(256), "190100");
assert_matches_reference(Value::Int(65535), "19FFFF");
assert_matches_reference(Value::Int(65536), "1A00010000");
}
#[test]
fn integers_negative_minimal_length() {
assert_matches_reference(Value::Int(-1), "20");
assert_matches_reference(Value::Int(-24), "37");
assert_matches_reference(Value::Int(-25), "3818");
assert_matches_reference(Value::Int(-256), "38FF");
}
#[test]
fn integer_out_of_range_is_rejected() {
assert_eq!(
encode(&Value::Int(u64::MAX as i128 + 1)),
Err(IntOutOfRange(u64::MAX as i128 + 1))
);
let floor = -(1i128 << 64);
assert!(encode(&Value::Int(floor)).is_ok());
assert!(encode(&Value::Int(floor - 1)).is_err());
}
#[test]
fn byte_strings() {
assert_matches_reference(Value::Bytes(vec![]), "40");
assert_matches_reference(Value::Bytes(b"hello".to_vec()), "4568656C6C6F");
}
#[test]
fn text_and_atom_equivalent_encoding() {
assert_matches_reference(Value::text("hello"), "6568656C6C6F");
assert_matches_reference(Value::text("true"), "6474727565");
}
#[test]
fn lists() {
assert_matches_reference(Value::List(vec![]), "80");
assert_matches_reference(
Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
"83010203",
);
}
#[test]
fn floats_always_binary64() {
assert_matches_reference(Value::Float(0.0), "FB0000000000000000");
assert_matches_reference(Value::Float(12345.6789), "FB40C81CD6E631F8A1");
}
#[test]
fn map_keys_sorted_by_encoded_bytes_not_input_order() {
assert_matches_reference(
Value::Map(vec![
(Value::text("b"), Value::Int(2)),
(Value::text("a"), Value::Int(1)),
]),
"A2616101616202",
);
}
#[test]
fn map_keys_sorted_lexicographically_same_length() {
assert_matches_reference(
Value::Map(vec![
(Value::text("zebra"), Value::Int(1)),
(Value::text("apple"), Value::Int(2)),
]),
"A2656170706C6502657A6562726101",
);
}
#[test]
fn map_keys_shorter_sorts_first_when_prefix() {
assert_matches_reference(
Value::Map(vec![
(Value::text("aa"), Value::Int(1)),
(Value::text("a"), Value::Int(2)),
(Value::text("aaa"), Value::Int(3)),
]),
"A3616102626161016361616103",
);
}
#[test]
fn null_alone() {
assert_matches_reference(Value::Null, "F6");
}
#[test]
fn nested_structure_with_null() {
assert_matches_reference(
Value::Map(vec![
(Value::text("name"), Value::text("macula")),
(
Value::text("nums"),
Value::List(vec![Value::Int(1), Value::Int(2), Value::Int(3)]),
),
(Value::text("nil"), Value::Null),
]),
"A3636E696CF6646E616D65666D6163756C61646E756D7383010203",
);
}
#[test]
fn frame_shaped_map() {
let node_id: Vec<u8> = (1u8..=32).collect();
assert_matches_reference(
Value::Map(vec![
(Value::text("node_id"), Value::Bytes(node_id)),
(Value::text("version"), Value::Int(2)),
(Value::text("frame_type"), Value::text("connect")),
(Value::text("capabilities"), Value::Int(0)),
]),
"A4676E6F64655F696458200102030405060708090A0B0C0D0E0F101112131415161718191A1B1C1D1E1F206776657273696F6E026A6672616D655F7479706567636F6E6E6563746C6361706162696C697469657300",
);
}
#[test]
fn decode_rejects_tags() {
assert_eq!(decode(&[0xC0]), Err(DecodeError::UnsupportedMajorType(6)));
}
#[test]
fn decode_rejects_trailing_bytes() {
assert_eq!(decode(&[0x00, 0xFF]), Err(DecodeError::TrailingBytes));
}
#[test]
fn decode_rejects_truncated_input() {
assert_eq!(decode(&[0x18]), Err(DecodeError::Truncated));
}
fn nested_list_payload(depth: usize) -> Vec<u8> {
let mut buf = vec![0x81u8; depth];
buf.push(0x00);
buf
}
#[test]
fn decode_accepts_nesting_at_the_depth_limit() {
let bytes = nested_list_payload(MAX_NESTING_DEPTH);
assert!(decode(&bytes).is_ok());
}
#[test]
fn decode_rejects_nesting_one_past_the_depth_limit() {
let bytes = nested_list_payload(MAX_NESTING_DEPTH + 1);
assert_eq!(decode(&bytes), Err(DecodeError::NestingTooDeep));
}
#[test]
fn decode_rejects_extreme_nesting_without_crashing() {
let bytes = nested_list_payload(100_000);
assert_eq!(decode(&bytes), Err(DecodeError::NestingTooDeep));
}
#[test]
fn decode_duplicate_map_keys_last_write_wins() {
let bytes = hex("A2616101616102");
let decoded = decode(&bytes).expect("valid map");
match decoded {
Value::Map(pairs) => {
assert_eq!(pairs.len(), 1);
assert_eq!(pairs[0], (Value::text("a"), Value::Int(2)));
}
other => panic!("expected a map, got {other:?}"),
}
}
#[test]
fn decode_duplicate_map_key_overwrites_its_original_slot_not_the_end() {
let map = Value::Map(vec![
(Value::text("a"), Value::Int(1)),
(Value::text("b"), Value::Int(2)),
(Value::text("c"), Value::Int(3)),
]);
let mut bytes = encode(&map).expect("encodable");
assert_eq!(bytes[0] & 0x1F, 3, "expected a 3-entry map header");
bytes[0] = (bytes[0] & 0xE0) | 4;
bytes.extend_from_slice(&encode(&Value::text("b")).unwrap());
bytes.extend_from_slice(&encode(&Value::Int(99)).unwrap());
let decoded = decode(&bytes).expect("valid map");
match decoded {
Value::Map(pairs) => {
assert_eq!(
pairs,
vec![
(Value::text("a"), Value::Int(1)),
(Value::text("b"), Value::Int(99)),
(Value::text("c"), Value::Int(3)),
]
);
}
other => panic!("expected a map, got {other:?}"),
}
}
#[test]
fn decode_map_with_many_distinct_keys_is_not_quadratic() {
let n: i128 = 20_000;
let pairs: Vec<(Value, Value)> = (0..n).map(|i| (Value::Int(i), Value::Int(0))).collect();
let bytes = encode(&Value::Map(pairs)).expect("encodable");
let start = std::time::Instant::now();
let decoded = decode(&bytes).expect("valid map");
let elapsed = start.elapsed();
match decoded {
Value::Map(decoded_pairs) => assert_eq!(decoded_pairs.len(), n as usize),
other => panic!("expected a map, got {other:?}"),
}
assert!(
elapsed < std::time::Duration::from_secs(2),
"decoding {n} distinct-keyed entries took {elapsed:?} -- \
looks like decode_map regressed to O(n^2)"
);
}
#[test]
fn decode_map_with_a_large_deeply_nested_key_is_not_quadratic_in_depth() {
let blob_len = 512 * 1024;
let mut bytes = vec![0xA1u8; MAX_NESTING_DEPTH];
bytes.push(0x5A); bytes.extend_from_slice(&(blob_len as u32).to_be_bytes());
bytes.extend(std::iter::repeat_n(0x41u8, blob_len));
bytes.extend(std::iter::repeat_n(0x00u8, MAX_NESTING_DEPTH));
let start = std::time::Instant::now();
let decoded = decode(&bytes).expect("valid, maximally-nested map-key chain");
let elapsed = start.elapsed();
let mut cursor = &decoded;
for _ in 0..MAX_NESTING_DEPTH {
match cursor {
Value::Map(pairs) if pairs.len() == 1 => cursor = &pairs[0].0,
other => panic!("expected a 1-entry map at this nesting level, got {other:?}"),
}
}
match cursor {
Value::Bytes(b) => assert_eq!(b.len(), blob_len),
other => panic!("expected the innermost key to be Bytes, got {other:?}"),
}
assert!(
elapsed < std::time::Duration::from_secs(5),
"decoding a {MAX_NESTING_DEPTH}-deep map-key chain around a {blob_len}-byte blob \
took {elapsed:?} -- looks like key canonicalization regressed to re-deriving a \
key's bytes at every ancestor level instead of reusing decode_one's own"
);
}
#[test]
fn get_finds_a_field_by_text_key() {
let map = Value::Map(vec![(Value::text("a"), Value::Int(1))]);
assert_eq!(map.get("a"), Some(&Value::Int(1)));
assert_eq!(map.get("missing"), None);
}
#[test]
fn get_on_a_non_map_is_none() {
assert_eq!(Value::Int(1).get("a"), None);
}
#[test]
fn without_removes_only_the_named_keys() {
let map = Value::Map(vec![
(Value::text("a"), Value::Int(1)),
(Value::text("b"), Value::Int(2)),
(Value::text("c"), Value::Int(3)),
]);
let stripped = map.without(&["b"]);
assert_eq!(stripped.get("a"), Some(&Value::Int(1)));
assert_eq!(stripped.get("b"), None);
assert_eq!(stripped.get("c"), Some(&Value::Int(3)));
}
#[test]
fn with_field_replaces_an_existing_key_in_place() {
let map =
Value::Map(vec![(Value::text("a"), Value::Int(1))]).with_field("a", Value::Int(2));
assert_eq!(map.get("a"), Some(&Value::Int(2)));
match map {
Value::Map(pairs) => assert_eq!(pairs.len(), 1),
_ => panic!("expected a map"),
}
}
#[test]
fn with_field_appends_a_new_key() {
let map = Value::Map(vec![]).with_field("a", Value::Int(1));
assert_eq!(map.get("a"), Some(&Value::Int(1)));
}
}