use crate::{
CborError, CborValue, MAX_DAG_CBOR_CONTAINER_ENTRIES, MAX_DAG_CBOR_INPUT_LEN,
MAX_DAG_CBOR_NODES, MAX_NESTING_DEPTH,
};
use std::cmp::Ordering;
use std::str;
use zeroize::Zeroizing;
const MT_UINT: u8 = 0;
const MT_NEGINT: u8 = 1;
const MT_BYTES: u8 = 2;
const MT_STRING: u8 = 3;
const MT_ARRAY: u8 = 4;
const MT_MAP: u8 = 5;
const MIN_ELEMENT_ENCODED_LEN: usize = 1;
#[derive(Default)]
struct DecodeBudget {
seen: usize,
pending: usize,
}
impl DecodeBudget {
fn enter(&mut self) -> Result<(), CborError> {
if self.seen != 0 {
self.pending = self
.pending
.checked_sub(1)
.ok_or(CborError::NodeLimitExceeded)?;
}
self.seen = self
.seen
.checked_add(1)
.ok_or(CborError::NodeLimitExceeded)?;
if self.seen > MAX_DAG_CBOR_NODES {
return Err(CborError::NodeLimitExceeded);
}
Ok(())
}
fn reserve_children(&mut self, count: usize) -> Result<(), CborError> {
let pending = self
.pending
.checked_add(count)
.ok_or(CborError::NodeLimitExceeded)?;
let total = self
.seen
.checked_add(pending)
.ok_or(CborError::NodeLimitExceeded)?;
if total > MAX_DAG_CBOR_NODES {
return Err(CborError::NodeLimitExceeded);
}
self.pending = pending;
Ok(())
}
}
pub fn decode_dag_cbor(bytes: &[u8]) -> Result<CborValue, CborError> {
if bytes.len() > MAX_DAG_CBOR_INPUT_LEN {
return Err(CborError::InputTooLarge);
}
let mut budget = DecodeBudget::default();
let (value, offset) = decode_value(bytes, 0, 0, &mut budget)?;
let mut value = Zeroizing::new(value);
if offset != bytes.len() {
return Err(CborError::TrailingBytes);
}
Ok(core::mem::replace(&mut *value, CborValue::Null))
}
fn decode_value(
bytes: &[u8],
mut offset: usize,
depth: usize,
budget: &mut DecodeBudget,
) -> Result<(CborValue, usize), CborError> {
if offset >= bytes.len() {
return Err(CborError::UnexpectedEnd);
}
budget.enter()?;
let first = bytes[offset];
offset = offset.checked_add(1).ok_or(CborError::OffsetOverflow)?;
let major = first >> 5;
let ai = first & 0x1f;
if major == 7 && (25..=27).contains(&ai) {
return Err(CborError::DisallowedSimpleValue {
value: u64::from(ai),
});
}
let (arg, new_offset) = read_argument(bytes, offset, ai)?;
offset = new_offset;
match major {
MT_UINT => Ok((
CborValue::Int(i64::try_from(arg).map_err(|_| CborError::IntegerOutOfRange)?),
offset,
)),
MT_NEGINT => {
let magnitude = i128::from(arg);
let value = (-1_i128)
.checked_sub(magnitude)
.ok_or(CborError::IntegerOutOfRange)?;
Ok((
CborValue::Int(i64::try_from(value).map_err(|_| CborError::IntegerOutOfRange)?),
offset,
))
}
MT_BYTES => {
let (b, off) = extract_bytes(bytes, offset, arg)?;
Ok((CborValue::Bytes(b), off))
}
MT_STRING => {
let (s, off) = extract_string(bytes, offset, arg)?;
Ok((CborValue::String(s), off))
}
MT_ARRAY => {
let child_depth = descend(depth)?;
let item_count = usize::try_from(arg).map_err(|_| CborError::LengthTooLarge)?;
bounded_capacity(item_count, bytes.len(), offset)?;
if item_count > MAX_DAG_CBOR_CONTAINER_ENTRIES {
return Err(CborError::ContainerEntriesExceeded);
}
budget.reserve_children(item_count)?;
let mut items = Zeroizing::new(Vec::with_capacity(item_count));
let mut off = offset;
for _ in 0..item_count {
let (v, next) = decode_value(bytes, off, child_depth, budget)?;
items.push(v);
off = next;
}
Ok((CborValue::Array(core::mem::take(&mut *items)), off))
}
MT_MAP => {
let child_depth = descend(depth)?;
let entry_count = usize::try_from(arg).map_err(|_| CborError::LengthTooLarge)?;
let entry_min = MIN_ELEMENT_ENCODED_LEN
.checked_mul(2)
.ok_or(CborError::OffsetOverflow)?;
bounded_capacity_with_min(entry_count, bytes.len(), offset, entry_min)?;
if entry_count > MAX_DAG_CBOR_CONTAINER_ENTRIES {
return Err(CborError::ContainerEntriesExceeded);
}
let child_count = entry_count
.checked_mul(2)
.ok_or(CborError::NodeLimitExceeded)?;
budget.reserve_children(child_count)?;
let mut entries = Zeroizing::new(Vec::with_capacity(entry_count));
let mut off = offset;
let mut last_key_bytes: Option<Zeroizing<Vec<u8>>> = None;
for _ in 0..entry_count {
let (key_val, key_off) = decode_value(bytes, off, child_depth, budget)?;
off = key_off;
let mut key_val = Zeroizing::new(key_val);
let mut key = Zeroizing::new(match &mut *key_val {
CborValue::String(s) => core::mem::take(s),
_ => return Err(CborError::MapKeyMustBeString),
});
let key_bytes = Zeroizing::new(key.as_bytes().to_vec());
if let Some(prev) = &last_key_bytes {
match compare_bytes(prev, &key_bytes) {
Ordering::Less => {}
Ordering::Equal => return Err(CborError::DuplicateMapKey),
Ordering::Greater => return Err(CborError::MapKeysOutOfOrder),
}
}
last_key_bytes = Some(key_bytes);
let (val, val_off) = decode_value(bytes, off, child_depth, budget)?;
off = val_off;
entries.push((core::mem::take(&mut *key), val));
}
Ok((CborValue::Map(core::mem::take(&mut *entries)), off))
}
7 => match arg {
20 => Ok((CborValue::Bool(false), offset)),
21 => Ok((CborValue::Bool(true), offset)),
22 => Ok((CborValue::Null, offset)),
_ => Err(CborError::DisallowedSimpleValue { value: arg }),
},
_ => Err(CborError::DisallowedMajorType { major }),
}
}
fn read_argument(bytes: &[u8], offset: usize, ai: u8) -> Result<(u64, usize), CborError> {
match ai {
n @ 0..=23 => Ok((u64::from(n), offset)),
24 => {
let end = checked_end(offset, 1)?;
if end > bytes.len() {
return Err(CborError::TruncatedArgument);
}
let value = u64::from(bytes[offset]);
if value < 24 {
return Err(CborError::NonCanonicalInteger);
}
Ok((value, end))
}
25 => {
let end = checked_end(offset, 2)?;
if end > bytes.len() {
return Err(CborError::TruncatedArgument);
}
let val = u16::from_be_bytes([bytes[offset], bytes[offset + 1]]);
if val < 256 {
return Err(CborError::NonCanonicalInteger);
}
Ok((u64::from(val), end))
}
26 => {
let end = checked_end(offset, 4)?;
if end > bytes.len() {
return Err(CborError::TruncatedArgument);
}
let val = u32::from_be_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
]);
if val < 65536 {
return Err(CborError::NonCanonicalInteger);
}
Ok((u64::from(val), end))
}
27 => {
let end = checked_end(offset, 8)?;
if end > bytes.len() {
return Err(CborError::TruncatedArgument);
}
let val = u64::from_be_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
bytes[offset + 4],
bytes[offset + 5],
bytes[offset + 6],
bytes[offset + 7],
]);
if val < 0x1_0000_0000 {
return Err(CborError::NonCanonicalInteger);
}
Ok((val, end))
}
_ => Err(CborError::UnsupportedAdditionalInfo),
}
}
fn extract_bytes(bytes: &[u8], offset: usize, len: u64) -> Result<(Vec<u8>, usize), CborError> {
let len = usize::try_from(len).map_err(|_| CborError::LengthTooLarge)?;
let end = checked_end(offset, len)?;
if end > bytes.len() {
return Err(CborError::TruncatedBytes);
}
Ok((bytes[offset..end].to_vec(), end))
}
fn extract_string(bytes: &[u8], offset: usize, len: u64) -> Result<(String, usize), CborError> {
let len = usize::try_from(len).map_err(|_| CborError::LengthTooLarge)?;
let end = checked_end(offset, len)?;
let raw = bytes.get(offset..end).ok_or(CborError::TruncatedBytes)?;
let s = str::from_utf8(raw).map_err(|_| CborError::InvalidUtf8)?;
Ok((s.to_string(), end))
}
fn checked_end(offset: usize, len: usize) -> Result<usize, CborError> {
offset.checked_add(len).ok_or(CborError::OffsetOverflow)
}
fn descend(depth: usize) -> Result<usize, CborError> {
let next = depth.checked_add(1).ok_or(CborError::OffsetOverflow)?;
if next > MAX_NESTING_DEPTH {
return Err(CborError::DepthExceeded);
}
Ok(next)
}
fn bounded_capacity(count: usize, total_len: usize, offset: usize) -> Result<usize, CborError> {
bounded_capacity_with_min(count, total_len, offset, MIN_ELEMENT_ENCODED_LEN)
}
fn bounded_capacity_with_min(
count: usize,
total_len: usize,
offset: usize,
min_element_len: usize,
) -> Result<usize, CborError> {
let remaining = total_len.saturating_sub(offset);
let max_possible = remaining / min_element_len.max(1);
if count > max_possible {
return Err(CborError::ContainerLengthExceedsInput);
}
Ok(count)
}
fn compare_bytes(a: &[u8], b: &[u8]) -> Ordering {
match a.len().cmp(&b.len()) {
Ordering::Equal => a.cmp(b),
ordering => ordering,
}
}