use serde::{Deserialize, Deserializer, de::DeserializeOwned};
use std::io::Cursor;
pub fn deserialize_present_option<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
where
D: Deserializer<'de>,
T: Deserialize<'de>,
{
Option::<T>::deserialize(deserializer)
}
pub fn from_slice_exact<T: DeserializeOwned>(
bytes: &[u8],
) -> Result<T, ciborium::de::Error<std::io::Error>> {
preflight(bytes)?;
let mut reader = Cursor::new(bytes);
let value = ciborium::from_reader(&mut reader)?;
let consumed = usize::try_from(reader.position()).unwrap_or(usize::MAX);
if consumed != bytes.len() {
return Err(ciborium::de::Error::semantic(
consumed,
"trailing bytes after CBOR value",
));
}
Ok(value)
}
fn preflight(bytes: &[u8]) -> Result<(), ciborium::de::Error<std::io::Error>> {
fn invalid() -> ciborium::de::Error<std::io::Error> {
ciborium::de::Error::semantic(0, "CBOR ledger size, nesting or structural bound exceeded")
}
fn value(
bytes: &[u8],
pos: &mut usize,
depth: usize,
) -> Result<(), ciborium::de::Error<std::io::Error>> {
if depth > crate::constants::MAX_LEDGER_NESTING {
return Err(invalid());
}
let head = *bytes.get(*pos).ok_or_else(invalid)?;
*pos += 1;
let major = head >> 5;
let info = head & 31;
let n = match info {
0..=23 => u64::from(info),
24..=27 => {
let width = 1_usize << (info - 24);
let end = pos.checked_add(width).ok_or_else(invalid)?;
let data = bytes.get(*pos..end).ok_or_else(invalid)?;
*pos = end;
data.iter().fold(0_u64, |n, b| (n << 8) | u64::from(*b))
}
_ => return Err(invalid()),
};
match major {
0 | 1 | 7 => (),
2 | 3 => {
let n = usize::try_from(n).map_err(|_| invalid())?;
if major == 2 && n > crate::constants::MAX_COMMITTED_PAYLOAD_BYTES {
return Err(invalid());
}
*pos = pos
.checked_add(n)
.filter(|end| *end <= bytes.len())
.ok_or_else(invalid)?;
}
4 | 5 => {
let count = if major == 5 {
n.checked_mul(2).ok_or_else(invalid)?
} else {
n
};
if count > (bytes.len() - *pos) as u64 {
return Err(invalid());
}
for _ in 0..count {
value(bytes, pos, depth + 1)?;
}
}
6 => value(bytes, pos, depth + 1)?,
_ => return Err(invalid()),
}
Ok(())
}
if bytes.len() > crate::constants::MAX_LEDGER_RECORD_BYTES {
return Err(invalid());
}
let mut pos = 0;
value(bytes, &mut pos, 0)?;
if pos != bytes.len() {
return Err(ciborium::de::Error::semantic(
pos,
"trailing bytes after CBOR value",
));
}
Ok(())
}
pub fn deserialize_records<'de, D, T>(deserializer: D) -> Result<Vec<T>, D::Error>
where
D: Deserializer<'de>,
T: Deserialize<'de>,
{
deserialize_bounded_vec::<D, T, 255>(deserializer)
}
pub fn deserialize_history<'de, D, T>(deserializer: D) -> Result<Vec<T>, D::Error>
where
D: Deserializer<'de>,
T: Deserialize<'de>,
{
deserialize_bounded_vec::<D, T, { crate::constants::MAX_LEDGER_GENERATIONS }>(deserializer)
}
pub fn deserialize_bounded_vec<'de, D, T, const LIMIT: usize>(
deserializer: D,
) -> Result<Vec<T>, D::Error>
where
D: Deserializer<'de>,
T: Deserialize<'de>,
{
struct Bounded<T, const LIMIT: usize>(std::marker::PhantomData<T>);
impl<'de, T: Deserialize<'de>, const LIMIT: usize> serde::de::Visitor<'de> for Bounded<T, LIMIT> {
type Value = Vec<T>;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(formatter, "at most {LIMIT} ledger entries")
}
fn visit_seq<A: serde::de::SeqAccess<'de>>(
self,
mut seq: A,
) -> Result<Self::Value, A::Error> {
if seq.size_hint().is_some_and(|n| n > LIMIT) {
return Err(serde::de::Error::custom("ledger collection limit exceeded"));
}
let mut values = Vec::new();
while let Some(value) = seq.next_element()? {
if values.len() == LIMIT {
return Err(serde::de::Error::custom("ledger collection limit exceeded"));
}
values.push(value);
}
Ok(values)
}
}
deserializer.deserialize_seq(Bounded::<T, LIMIT>(std::marker::PhantomData))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn oversized_byte_string_rejects_before_deserializer_runs() {
#[derive(Debug)]
struct MustNotDecode;
impl<'de> Deserialize<'de> for MustNotDecode {
fn deserialize<D: Deserializer<'de>>(_: D) -> Result<Self, D::Error> {
panic!("preflight must reject before serde can allocate");
}
}
let len = crate::constants::MAX_COMMITTED_PAYLOAD_BYTES + 1;
let mut bytes = vec![0x5a];
bytes.extend_from_slice(&u32::try_from(len).unwrap().to_be_bytes());
assert!(from_slice_exact::<MustNotDecode>(&bytes).is_err());
bytes.resize(5 + len, 0);
assert!(from_slice_exact::<MustNotDecode>(&bytes).is_err());
}
}