Skip to content

Commit 1995d1a

Browse files
authored
Merge branch 'main' into optimize-encode-buffer-reuse-and-cid-parse
2 parents 62f3e64 + 7970600 commit 1995d1a

3 files changed

Lines changed: 87 additions & 35 deletions

File tree

pytests/test_dag_cbor.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -277,3 +277,25 @@ def test_roundtrip_valid_cid_with_short_tag() -> None:
277277
encoded = libipld.encode_dag_cbor(decoded)
278278

279279
assert encoded == encoded_bytes
280+
281+
282+
def test_dag_cbor_decode_array_length_exceeds_data_error() -> None:
283+
# 9b0000000040000000 - array claiming 2**30 elements with no payload
284+
with pytest.raises(ValueError) as exc_info:
285+
libipld.decode_dag_cbor(bytes.fromhex('9b0000000040000000'))
286+
287+
assert 'Array length exceeds remaining data' in str(exc_info.value)
288+
289+
290+
def test_dag_cbor_decode_map_length_exceeds_data_error() -> None:
291+
# bb0000000040000000 - map claiming 2**30 entries with no payload
292+
with pytest.raises(ValueError) as exc_info:
293+
libipld.decode_dag_cbor(bytes.fromhex('bb0000000040000000'))
294+
295+
assert 'Map length exceeds remaining data' in str(exc_info.value)
296+
297+
298+
def test_dag_cbor_decode_error_mid_array() -> None:
299+
# 83010263616263 - [1, 2, "abc"] truncated after the second element's header
300+
with pytest.raises(ValueError):
301+
libipld.decode_dag_cbor(bytes.fromhex('8301026361'))

pytests/test_decode_dag_cbor_multi.py

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,3 +36,21 @@ def test_decode_dag_cbor_multi(data) -> None:
3636
decoded = libipld.decode_dag_cbor_multi(encoded)
3737
assert len(decoded) == len(objects)
3838
assert decoded == objects
39+
40+
41+
def test_decode_dag_cbor_multi_corrupt_trailing_data_error() -> None:
42+
encoded = libipld.encode_dag_cbor({'abc': 1})
43+
44+
with pytest.raises(ValueError) as exc_info:
45+
# 9b0000000040000000 - array claiming 2**30 elements with no payload
46+
libipld.decode_dag_cbor_multi(encoded + bytes.fromhex('9b0000000040000000'))
47+
48+
assert 'Failed to decode DAG-CBOR' in str(exc_info.value)
49+
50+
51+
def test_decode_dag_cbor_multi_truncated_object_error() -> None:
52+
encoded = libipld.encode_dag_cbor({'abc': 1})
53+
54+
with pytest.raises(ValueError):
55+
# second object is truncated mid-string
56+
libipld.decode_dag_cbor_multi(encoded + bytes.fromhex('6361'))

src/dag_cbor/de.rs

Lines changed: 47 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -65,12 +65,22 @@ where
6565
.into()
6666
}
6767
major::ARRAY => {
68-
let len: ffi::Py_ssize_t = types::Array::len(r)?
69-
.ok_or_else(|| anyhow!("Array must contain length"))?
70-
.try_into()?;
68+
let len = types::Array::len(r)?.ok_or_else(|| anyhow!("Array must contain length"))?;
69+
// Every element costs at least one byte; reject a claimed length
70+
// beyond the remaining input before allocating for it.
71+
if r.fill(len)?.as_ref().len() < len {
72+
return Err(anyhow!("Array length exceeds remaining data"));
73+
}
74+
let len: ffi::Py_ssize_t = len.try_into()?;
7175

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

7585
for i in 0..len {
7686
ffi::PyList_SET_ITEM(
@@ -80,12 +90,17 @@ where
8090
);
8191
}
8292

83-
let list: Bound<'_, PyList> = Bound::from_owned_ptr(py, ptr).cast_into_unchecked();
8493
list.into_pyobject(py)?.into()
8594
}
8695
}
8796
major::MAP => {
8897
let len = types::Map::len(r)?.ok_or_else(|| anyhow!("Map must contain length"))?;
98+
// Every entry costs at least two bytes (key + value); reject a
99+
// claimed length beyond the remaining input before presizing.
100+
let need = len.saturating_mul(2);
101+
if r.fill(need)?.as_ref().len() < need {
102+
return Err(anyhow!("Map length exceeds remaining data"));
103+
}
89104
// Length is known up front; presize to avoid rehashes as we fill.
90105
let dict = unsafe {
91106
let ptr = new_presized(len);
@@ -190,18 +205,28 @@ where
190205
})
191206
}
192207

208+
// Wrap a decode failure; an error already set on the interpreter (e.g. the
209+
// RecursionError `restore`d above) wins, with the decode error as its cause.
210+
fn decode_error(py: Python, e: anyhow::Error) -> PyErr {
211+
let err = value_error("Failed to decode DAG-CBOR", e.to_string());
212+
if let Some(py_err) = PyErr::take(py) {
213+
py_err.set_cause(py, Option::from(err));
214+
py_err
215+
} else {
216+
err
217+
}
218+
}
219+
193220
#[pyfunction]
194221
pub fn decode_dag_cbor_multi<'py>(py: Python<'py>, data: &[u8]) -> PyResult<Bound<'py, PyList>> {
195222
let mut reader = SliceReader::new(data);
196223
let decoded_parts = PyList::empty(py);
197224
let max_depth = current_recursion_limit();
198225

199-
loop {
200-
let py_object = to_pyobject(py, &mut reader, 0, max_depth);
201-
if let Ok(py_object) = py_object {
202-
decoded_parts.append(py_object)?;
203-
} else {
204-
break;
226+
while !reader.fill(1)?.as_ref().is_empty() {
227+
match to_pyobject(py, &mut reader, 0, max_depth) {
228+
Ok(py_object) => decoded_parts.append(py_object)?,
229+
Err(e) => return Err(decode_error(py, e)),
205230
}
206231
}
207232

@@ -212,31 +237,18 @@ pub fn decode_dag_cbor_multi<'py>(py: Python<'py>, data: &[u8]) -> PyResult<Boun
212237
pub fn decode_dag_cbor(py: Python, data: &[u8]) -> PyResult<Py<PyAny>> {
213238
let mut reader = SliceReader::new(data);
214239
let max_depth = current_recursion_limit();
215-
let py_object = to_pyobject(py, &mut reader, 0, max_depth);
216-
if let Ok(py_object) = py_object {
217-
// check for any remaining data in the reader
218-
if reader.fill(1)?.as_ref().is_empty() {
219-
Ok(py_object)
220-
} else {
221-
Err(value_error(
222-
"Failed to decode DAG-CBOR",
223-
"Invalid DAG-CBOR: contains multiple objects (CBOR sequence)".to_string(),
224-
))
225-
}
226-
} else {
227-
let err = value_error(
228-
"Failed to decode DAG-CBOR",
229-
py_object.unwrap_err().to_string(),
230-
);
231-
232-
if let Some(py_err) = PyErr::take(py) {
233-
py_err.set_cause(py, Option::from(err));
234-
// in case something set global interpreter’s error,
235-
// for example C FFI function, we should return it
236-
// the real case: RecursionError (set by Py_EnterRecursiveCall)
237-
Err(py_err)
238-
} else {
239-
Err(err)
240+
match to_pyobject(py, &mut reader, 0, max_depth) {
241+
Ok(py_object) => {
242+
// check for any remaining data in the reader
243+
if reader.fill(1)?.as_ref().is_empty() {
244+
Ok(py_object)
245+
} else {
246+
Err(value_error(
247+
"Failed to decode DAG-CBOR",
248+
"Invalid DAG-CBOR: contains multiple objects (CBOR sequence)".to_string(),
249+
))
250+
}
240251
}
252+
Err(e) => Err(decode_error(py, e)),
241253
}
242254
}

0 commit comments

Comments
 (0)