From 26daa74049d962241f72e4436085d4cda05a469c Mon Sep 17 00:00:00 2001 From: bit2swaz Date: Fri, 24 Jul 2026 00:23:52 +0530 Subject: [PATCH] fix(arrow-array): make ArrowArrayStreamReader::try_new unsafe --- arrow-array/src/ffi_stream.rs | 32 +++++++++++++++++++++++++++----- arrow-pyarrow/src/lib.rs | 6 ++++-- 2 files changed, 31 insertions(+), 7 deletions(-) diff --git a/arrow-array/src/ffi_stream.rs b/arrow-array/src/ffi_stream.rs index 9a09c3753dd4..cf4a5890ae9b 100644 --- a/arrow-array/src/ffi_stream.rs +++ b/arrow-array/src/ffi_stream.rs @@ -318,8 +318,27 @@ fn get_stream_schema(stream_ptr: *mut FFI_ArrowArrayStream) -> Result impl ArrowArrayStreamReader { /// Creates a new `ArrowArrayStreamReader` from a `FFI_ArrowArrayStream`. /// This is used to import from the C Stream Interface. - #[allow(dead_code)] - pub fn try_new(mut stream: FFI_ArrowArrayStream) -> Result { + /// + /// # Safety + /// + /// The provided `stream` must be a valid [`FFI_ArrowArrayStream`] per the + /// [C Stream Interface]: its callbacks (`get_schema`, `get_next`, `get_last_error`, + /// `release`) and `private_data` must be either null/`None` or valid to call for + /// the lifetime of the returned reader. Passing a stream with garbage callbacks + /// is undefined behavior. + /// + /// [C Stream Interface]: https://arrow.apache.org/docs/format/CStreamInterface.html + /// + /// # Example + /// + /// ```compile_fail + /// # use arrow_array::ffi_stream::{ArrowArrayStreamReader, FFI_ArrowArrayStream}; + /// # use arrow_array::RecordBatchReader; + /// // try_new is unsafe — this must not compile: + /// let stream = FFI_ArrowArrayStream::empty(); + /// let _reader = ArrowArrayStreamReader::try_new(stream); + /// ``` + pub unsafe fn try_new(mut stream: FFI_ArrowArrayStream) -> Result { if stream.release.is_none() { return Err(ArrowError::CDataInterface( "input stream is already released".to_string(), @@ -342,7 +361,8 @@ impl ArrowArrayStreamReader { /// /// See [`FFI_ArrowArrayStream::from_raw`] pub unsafe fn from_raw(raw_stream: *mut FFI_ArrowArrayStream) -> Result { - Self::try_new(unsafe { FFI_ArrowArrayStream::from_raw(raw_stream) }) + // SAFETY: caller upholds the same contract as try_new requires + unsafe { Self::try_new(FFI_ArrowArrayStream::from_raw(raw_stream)) } } /// Get the last error from `ArrowArrayStreamReader` @@ -488,7 +508,8 @@ mod tests { // Import through `FFI_ArrowArrayStream` as `ArrowArrayStreamReader` let stream = FFI_ArrowArrayStream::new(reader); - let stream_reader = ArrowArrayStreamReader::try_new(stream).unwrap(); + // SAFETY: stream was constructed via FFI_ArrowArrayStream::new, so callbacks are valid + let stream_reader = unsafe { ArrowArrayStreamReader::try_new(stream) }.unwrap(); let imported_schema = stream_reader.schema(); assert_eq!(imported_schema, schema); @@ -550,7 +571,8 @@ mod tests { // Import through `FFI_ArrowArrayStream` as `ArrowArrayStreamReader` let stream = FFI_ArrowArrayStream::new(reader); - let stream_reader = ArrowArrayStreamReader::try_new(stream).unwrap(); + // SAFETY: stream was constructed via FFI_ArrowArrayStream::new, so callbacks are valid + let stream_reader = unsafe { ArrowArrayStreamReader::try_new(stream) }.unwrap(); let imported_schema = stream_reader.schema(); assert_eq!(imported_schema, schema); diff --git a/arrow-pyarrow/src/lib.rs b/arrow-pyarrow/src/lib.rs index cfe51b807aef..cb4f2e413083 100644 --- a/arrow-pyarrow/src/lib.rs +++ b/arrow-pyarrow/src/lib.rs @@ -384,7 +384,8 @@ impl FromPyArrow for ArrowArrayStreamReader { extract_capsule(&capsule, c"arrow_array_stream", "__arrow_c_stream__")?; let stream = unsafe { FFI_ArrowArrayStream::from_raw(stream_ptr.as_ptr()) }; - let stream_reader = ArrowArrayStreamReader::try_new(stream) + // SAFETY: stream came from a valid capsule pointer provided by the C stream interface + let stream_reader = unsafe { ArrowArrayStreamReader::try_new(stream) } .map_err(|err| PyValueError::new_err(err.to_string()))?; return Ok(stream_reader); @@ -403,7 +404,8 @@ impl FromPyArrow for ArrowArrayStreamReader { (&raw mut stream as Py_uintptr_t,), )?; - ArrowArrayStreamReader::try_new(stream) + // SAFETY: stream was populated by PyArrow's _export_to_c, which fills a valid C stream + unsafe { ArrowArrayStreamReader::try_new(stream) } .map_err(|err| PyValueError::new_err(err.to_string())) } }