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 {
TrailingBytes,
BadKey,
DuplicateKey,
InvalidText,
NestingTooDeep,
IntegerOutOfRange,
TooManyElements,
Malformed,
}
impl fmt::Display for DecodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
DecodeError::TrailingBytes => "bytes after the top-level value",
DecodeError::BadKey => "a map key that is neither text nor an integer",
DecodeError::DuplicateKey => "a duplicate map key",
DecodeError::InvalidText => "text that is not valid UTF-8",
DecodeError::NestingTooDeep => "arrays and maps nested more than 64 levels",
DecodeError::IntegerOutOfRange => "an integer below -2^63 or above 2^63-1",
DecodeError::TooManyElements => "more than 131072 items",
DecodeError::Malformed => "malformed",
})
}
}
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 = 64;
pub const MAX_ELEMENTS: usize = 131_072;
pub fn decode(bytes: &[u8]) -> Result<Value, DecodeError> {
let mut decoder = Decoder {
data: bytes,
pos: 0,
budget: MAX_ELEMENTS,
};
let value = decoder.item(0)?;
if decoder.pos != bytes.len() {
return Err(DecodeError::TrailingBytes);
}
Ok(value)
}
struct Decoder<'a> {
data: &'a [u8],
pos: usize,
budget: usize,
}
#[derive(PartialEq, Eq, Hash)]
enum KeyId {
Text(String),
Int(i128),
}
const MAX_SIZE_HINT: usize = 4;
impl Decoder<'_> {
fn item(&mut self, depth: usize) -> Result<Value, DecodeError> {
let head = self.take(1)?[0];
let (major, ai) = (head >> 5, head & 0x1F);
if major == 7 {
return self.simple_or_float(ai);
}
let arg = self.argument(ai)?;
self.count()?;
match major {
0 => integer(i128::from(arg), arg),
1 => integer(-1 - i128::from(arg), arg),
2 => Ok(Value::Bytes(self.take(arg)?.to_vec())),
3 => {
let bytes = self.take(arg)?;
std::str::from_utf8(bytes)
.map(|text| Value::Text(text.to_owned()))
.map_err(|_| DecodeError::InvalidText)
}
4 => self.list(arg, depth),
5 => self.map(arg, depth),
_ => Err(DecodeError::Malformed),
}
}
fn count(&mut self) -> Result<(), DecodeError> {
if self.budget == 0 {
return Err(DecodeError::TooManyElements);
}
self.budget -= 1;
Ok(())
}
fn take(&mut self, n: u64) -> Result<&[u8], DecodeError> {
let remaining = (self.data.len() - self.pos) as u64;
if n > remaining {
return Err(DecodeError::Malformed);
}
let start = self.pos;
self.pos += n as usize;
Ok(&self.data[start..self.pos])
}
fn argument(&mut self, ai: u8) -> Result<u64, DecodeError> {
let width = match ai {
0..=23 => return Ok(u64::from(ai)),
24 => 1,
25 => 2,
26 => 4,
27 => 8,
_ => return Err(DecodeError::Malformed),
};
Ok(self
.take(width)?
.iter()
.fold(0u64, |arg, &b| (arg << 8) | u64::from(b)))
}
fn simple_or_float(&mut self, ai: u8) -> Result<Value, DecodeError> {
match ai {
22 => {
self.count()?;
Ok(Value::Null)
}
25..=27 => {
let width = match ai {
25 => 2,
26 => 4,
_ => 8,
};
let bytes = self.take(width)?;
let value = match bytes.len() {
2 => half_to_f64(u16::from_be_bytes([bytes[0], bytes[1]])),
4 => f64::from(f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])),
_ => f64::from_be_bytes(bytes.try_into().map_err(|_| DecodeError::Malformed)?),
};
self.count()?;
if value.is_finite() {
Ok(Value::Float(value))
} else {
Err(DecodeError::Malformed)
}
}
0..=24 => {
self.argument(ai)?;
self.count()?;
Err(DecodeError::Malformed)
}
_ => Err(DecodeError::Malformed),
}
}
fn size_hint(&self, count: u64, items_per_element: usize) -> usize {
let bytes_left = (self.data.len() - self.pos) / items_per_element;
let budget_left = self.budget / items_per_element;
count
.min(bytes_left as u64)
.min(budget_left as u64)
.min(MAX_SIZE_HINT as u64) as usize
}
fn list(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
if depth >= MAX_NESTING_DEPTH {
return Err(DecodeError::NestingTooDeep);
}
let mut items = Vec::with_capacity(self.size_hint(count, 1));
for _ in 0..count {
items.push(self.item(depth + 1)?);
}
Ok(Value::List(items))
}
fn map(&mut self, count: u64, depth: usize) -> Result<Value, DecodeError> {
if depth >= MAX_NESTING_DEPTH {
return Err(DecodeError::NestingTooDeep);
}
let hint = self.size_hint(count, 2);
let mut pairs = Vec::with_capacity(hint);
let mut seen = std::collections::HashSet::with_capacity(hint);
for _ in 0..count {
let key = self.item(depth + 1)?;
let value = self.item(depth + 1)?;
let id = match &key {
Value::Text(text) => KeyId::Text(text.clone()),
Value::Int(n) => KeyId::Int(*n),
_ => return Err(DecodeError::BadKey),
};
if !seen.insert(id) {
return Err(DecodeError::DuplicateKey);
}
pairs.push((key, value));
}
Ok(Value::Map(pairs))
}
}
fn integer(value: i128, arg: u64) -> Result<Value, DecodeError> {
if arg >= 1 << 63 {
return Err(DecodeError::IntegerOutOfRange);
}
Ok(Value::Int(value))
}
fn half_to_f64(half: u16) -> f64 {
let sign = if half >> 15 == 1 { -1.0 } else { 1.0 };
let exp = (half >> 10) & 0x1F;
let frac = f64::from(half & 0x3FF);
match exp {
0 => sign * 2f64.powi(-24) * frac,
31 if frac == 0.0 => sign * f64::INFINITY,
31 => f64::NAN,
_ => sign * 2f64.powi(i32::from(exp) - 15) * (1.0 + frac / 1024.0),
}
}
#[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_trailing_bytes() {
assert_eq!(decode(&[0x00, 0xFF]), Err(DecodeError::TrailingBytes));
}
#[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 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)));
}
}