diff --git a/src/ffi/body.rs b/src/ffi/body.rs index e5a09e577b..04e45655ba 100644 --- a/src/ffi/body.rs +++ b/src/ffi/body.rs @@ -257,11 +257,10 @@ ffi_fn! { /// To avoid a memory leak, the copy must eventually be consumed by /// `hyper_buf_free`. /// - /// This returns `NULL` if allocating a new buffer fails. + /// This returns `NULL` if `buf` is `NULL`, or if allocating a new buffer + /// fails. fn hyper_buf_copy(buf: *const u8, len: size_t) -> *mut hyper_buf { - let slice = unsafe { - std::slice::from_raw_parts(buf, len) - }; + let slice = non_null!(buf, std::slice::from_raw_parts(buf, len), ptr::null_mut()); Box::into_raw(Box::new(hyper_buf(Bytes::copy_from_slice(slice)))) } ?= ptr::null_mut() } @@ -275,14 +274,14 @@ ffi_fn! { /// This pointer is borrowed data, and not valid once the `hyper_buf` is /// consumed/freed. fn hyper_buf_bytes(buf: *const hyper_buf) -> *const u8 { - unsafe { (*buf).0.as_ptr() } + non_null!(&*buf ?= ptr::null()).0.as_ptr() } ?= ptr::null() } ffi_fn! { /// Get the length of the bytes this buffer contains. fn hyper_buf_len(buf: *const hyper_buf) -> size_t { - unsafe { (*buf).0.len() } + non_null!(&*buf ?= 0).0.len() } } @@ -291,7 +290,7 @@ ffi_fn! { /// /// This should be used for any buffer once it is no longer needed. fn hyper_buf_free(buf: *mut hyper_buf) { - drop(unsafe { Box::from_raw(buf) }); + drop(non_null!(Box::from_raw(buf) ?= ())); } } @@ -300,3 +299,31 @@ unsafe impl AsTaskType for hyper_buf { hyper_task_return_type::HYPER_TASK_BUF } } + +#[cfg(test)] +mod tests { + use super::*; + + // `non_null!` also `debug_assert!`s, and a panic inside an `ffi_fn!` + // without a default aborts, so NULL handling can only be exercised + // without debug assertions (e.g. `cargo test --release`). + #[cfg(not(debug_assertions))] + #[test] + fn buf_functions_reject_null() { + assert!(hyper_buf_copy(ptr::null(), 5).is_null()); + assert!(hyper_buf_bytes(ptr::null()).is_null()); + assert_eq!(hyper_buf_len(ptr::null()), 0); + hyper_buf_free(ptr::null_mut()); + } + + #[test] + fn buf_copy_bytes_len() { + let src = b"hello"; + let buf = hyper_buf_copy(src.as_ptr(), src.len()); + assert!(!buf.is_null()); + assert_eq!(hyper_buf_len(buf), src.len()); + let bytes = unsafe { std::slice::from_raw_parts(hyper_buf_bytes(buf), hyper_buf_len(buf)) }; + assert_eq!(bytes, src); + hyper_buf_free(buf); + } +}