diff --git a/pytests/test_dag_cbor.py b/pytests/test_dag_cbor.py index 5c2bff1..76679e3 100644 --- a/pytests/test_dag_cbor.py +++ b/pytests/test_dag_cbor.py @@ -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')) diff --git a/pytests/test_decode_dag_cbor_multi.py b/pytests/test_decode_dag_cbor_multi.py index 8468883..8e80a1e 100644 --- a/pytests/test_decode_dag_cbor_multi.py +++ b/pytests/test_decode_dag_cbor_multi.py @@ -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')) diff --git a/src/dag_cbor/de.rs b/src/dag_cbor/de.rs index 5522ac7..07e9e0a 100644 --- a/src/dag_cbor/de.rs +++ b/src/dag_cbor/de.rs @@ -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( @@ -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); @@ -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> { 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)), } } @@ -211,31 +236,18 @@ pub fn decode_dag_cbor_multi<'py>(py: Python<'py>, data: &[u8]) -> PyResult PyResult> { 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)), } }