Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions pytests/test_dag_cbor.py
Original file line number Diff line number Diff line change
Expand Up @@ -277,3 +277,25 @@ def test_roundtrip_valid_cid_with_short_tag() -> None:
encoded = libipld.encode_dag_cbor(decoded)

assert encoded == encoded_bytes


def test_dag_cbor_decode_array_length_exceeds_data_error() -> None:
# 9b0000000040000000 - array claiming 2**30 elements with no payload
with pytest.raises(ValueError) as exc_info:
libipld.decode_dag_cbor(bytes.fromhex('9b0000000040000000'))

assert 'Array length exceeds remaining data' in str(exc_info.value)


def test_dag_cbor_decode_map_length_exceeds_data_error() -> None:
# bb0000000040000000 - map claiming 2**30 entries with no payload
with pytest.raises(ValueError) as exc_info:
libipld.decode_dag_cbor(bytes.fromhex('bb0000000040000000'))

assert 'Map length exceeds remaining data' in str(exc_info.value)


def test_dag_cbor_decode_error_mid_array() -> None:
# 83010263616263 - [1, 2, "abc"] truncated after the second element's header
with pytest.raises(ValueError):
libipld.decode_dag_cbor(bytes.fromhex('8301026361'))
18 changes: 18 additions & 0 deletions pytests/test_decode_dag_cbor_multi.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,3 +36,21 @@ def test_decode_dag_cbor_multi(data) -> None:
decoded = libipld.decode_dag_cbor_multi(encoded)
assert len(decoded) == len(objects)
assert decoded == objects


def test_decode_dag_cbor_multi_corrupt_trailing_data_error() -> None:
encoded = libipld.encode_dag_cbor({'abc': 1})

with pytest.raises(ValueError) as exc_info:
# 9b0000000040000000 - array claiming 2**30 elements with no payload
libipld.decode_dag_cbor_multi(encoded + bytes.fromhex('9b0000000040000000'))

assert 'Failed to decode DAG-CBOR' in str(exc_info.value)


def test_decode_dag_cbor_multi_truncated_object_error() -> None:
encoded = libipld.encode_dag_cbor({'abc': 1})

with pytest.raises(ValueError):
# second object is truncated mid-string
libipld.decode_dag_cbor_multi(encoded + bytes.fromhex('6361'))
82 changes: 47 additions & 35 deletions src/dag_cbor/de.rs
Original file line number Diff line number Diff line change
Expand Up @@ -64,12 +64,22 @@ where
.into()
}
major::ARRAY => {
let len: ffi::Py_ssize_t = types::Array::len(r)?
.ok_or_else(|| anyhow!("Array must contain length"))?
.try_into()?;
let len = types::Array::len(r)?.ok_or_else(|| anyhow!("Array must contain length"))?;
// Every element costs at least one byte; reject a claimed length
// beyond the remaining input before allocating for it.
if r.fill(len)?.as_ref().len() < len {
return Err(anyhow!("Array length exceeds remaining data"));
}
let len: ffi::Py_ssize_t = len.try_into()?;

unsafe {
let ptr = ffi::PyList_New(len);
if ptr.is_null() {
return Err(anyhow!(PyErr::fetch(py)));
}
// Owned before filling so an error mid-fill releases the list;
// list dealloc tolerates the remaining NULL slots.
let list: Bound<'_, PyList> = Bound::from_owned_ptr(py, ptr).cast_into_unchecked();

for i in 0..len {
ffi::PyList_SET_ITEM(
Expand All @@ -79,12 +89,17 @@ where
);
}

let list: Bound<'_, PyList> = Bound::from_owned_ptr(py, ptr).cast_into_unchecked();
list.into_pyobject(py)?.into()
}
}
major::MAP => {
let len = types::Map::len(r)?.ok_or_else(|| anyhow!("Map must contain length"))?;
// Every entry costs at least two bytes (key + value); reject a
// claimed length beyond the remaining input before presizing.
let need = len.saturating_mul(2);
if r.fill(need)?.as_ref().len() < need {
return Err(anyhow!("Map length exceeds remaining data"));
}
// Length is known up front; presize to avoid rehashes as we fill.
let dict = unsafe {
let ptr = new_presized(len);
Expand Down Expand Up @@ -189,18 +204,28 @@ where
})
}

// Wrap a decode failure; an error already set on the interpreter (e.g. the
// RecursionError `restore`d above) wins, with the decode error as its cause.
fn decode_error(py: Python, e: anyhow::Error) -> PyErr {
let err = value_error("Failed to decode DAG-CBOR", e.to_string());
if let Some(py_err) = PyErr::take(py) {
py_err.set_cause(py, Option::from(err));
py_err
} else {
err
}
}

#[pyfunction]
pub fn decode_dag_cbor_multi<'py>(py: Python<'py>, data: &[u8]) -> PyResult<Bound<'py, PyList>> {
let mut reader = SliceReader::new(data);
let decoded_parts = PyList::empty(py);
let max_depth = current_recursion_limit();

loop {
let py_object = to_pyobject(py, &mut reader, 0, max_depth);
if let Ok(py_object) = py_object {
decoded_parts.append(py_object)?;
} else {
break;
while !reader.fill(1)?.as_ref().is_empty() {
match to_pyobject(py, &mut reader, 0, max_depth) {
Ok(py_object) => decoded_parts.append(py_object)?,
Err(e) => return Err(decode_error(py, e)),
}
}

Expand All @@ -211,31 +236,18 @@ pub fn decode_dag_cbor_multi<'py>(py: Python<'py>, data: &[u8]) -> PyResult<Boun
pub fn decode_dag_cbor(py: Python, data: &[u8]) -> PyResult<Py<PyAny>> {
let mut reader = SliceReader::new(data);
let max_depth = current_recursion_limit();
let py_object = to_pyobject(py, &mut reader, 0, max_depth);
if let Ok(py_object) = py_object {
// check for any remaining data in the reader
if reader.fill(1)?.as_ref().is_empty() {
Ok(py_object)
} else {
Err(value_error(
"Failed to decode DAG-CBOR",
"Invalid DAG-CBOR: contains multiple objects (CBOR sequence)".to_string(),
))
}
} else {
let err = value_error(
"Failed to decode DAG-CBOR",
py_object.unwrap_err().to_string(),
);

if let Some(py_err) = PyErr::take(py) {
py_err.set_cause(py, Option::from(err));
// in case something set global interpreter’s error,
// for example C FFI function, we should return it
// the real case: RecursionError (set by Py_EnterRecursiveCall)
Err(py_err)
} else {
Err(err)
match to_pyobject(py, &mut reader, 0, max_depth) {
Ok(py_object) => {
// check for any remaining data in the reader
if reader.fill(1)?.as_ref().is_empty() {
Ok(py_object)
} else {
Err(value_error(
"Failed to decode DAG-CBOR",
"Invalid DAG-CBOR: contains multiple objects (CBOR sequence)".to_string(),
))
}
}
Err(e) => Err(decode_error(py, e)),
}
}
Loading