use arrow_schema::DataType;
use std::ffi::CStr;
use std::ptr::addr_of;
use std::{
ffi::CString,
os::raw::{c_char, c_int, c_void},
sync::Arc,
};
use arrow_data::ffi::FFI_ArrowArray;
use arrow_schema::{ArrowError, Schema, SchemaRef, ffi::FFI_ArrowSchema};
use crate::RecordBatchOptions;
use crate::array::Array;
use crate::array::StructArray;
use crate::ffi::from_ffi_and_data_type;
use crate::record_batch::{RecordBatch, RecordBatchReader};
type Result<T> = std::result::Result<T, ArrowError>;
#[cfg(not(all(target_family = "wasm", target_os = "unknown")))]
use libc::{EINVAL, EIO, ENOMEM, ENOSYS};
#[cfg(all(target_family = "wasm", target_os = "unknown"))]
const ENOMEM: i32 = 12;
#[cfg(all(target_family = "wasm", target_os = "unknown"))]
const EIO: i32 = 5;
#[cfg(all(target_family = "wasm", target_os = "unknown"))]
const EINVAL: i32 = 22;
#[cfg(all(target_family = "wasm", target_os = "unknown"))]
const ENOSYS: i32 = 38;
#[repr(C)]
#[derive(Debug)]
pub struct FFI_ArrowArrayStream {
get_schema: Option<unsafe extern "C" fn(arg1: *mut Self, out: *mut FFI_ArrowSchema) -> c_int>,
get_next: Option<unsafe extern "C" fn(arg1: *mut Self, out: *mut FFI_ArrowArray) -> c_int>,
get_last_error: Option<unsafe extern "C" fn(arg1: *mut Self) -> *const c_char>,
release: Option<unsafe extern "C" fn(arg1: *mut Self)>,
private_data: *mut c_void,
}
unsafe impl Send for FFI_ArrowArrayStream {}
unsafe extern "C" fn release_stream(stream: *mut FFI_ArrowArrayStream) {
if stream.is_null() {
return;
}
let stream = unsafe { &mut *stream };
stream.get_schema = None;
stream.get_next = None;
stream.get_last_error = None;
let private_data = unsafe { Box::from_raw(stream.private_data.cast::<StreamPrivateData>()) };
drop(private_data);
stream.release = None;
}
struct StreamPrivateData {
batch_reader: Box<dyn RecordBatchReader + Send>,
last_error: Option<CString>,
}
unsafe extern "C" fn get_schema(
stream: *mut FFI_ArrowArrayStream,
schema: *mut FFI_ArrowSchema,
) -> c_int {
ExportedArrayStream { stream }.get_schema(schema)
}
unsafe extern "C" fn get_next(
stream: *mut FFI_ArrowArrayStream,
array: *mut FFI_ArrowArray,
) -> c_int {
ExportedArrayStream { stream }.get_next(array)
}
unsafe extern "C" fn get_last_error(stream: *mut FFI_ArrowArrayStream) -> *const c_char {
let mut ffi_stream = ExportedArrayStream { stream };
match ffi_stream.get_last_error() {
Some(err_string) => err_string.as_ptr(),
None => std::ptr::null(),
}
}
impl Drop for FFI_ArrowArrayStream {
fn drop(&mut self) {
match self.release {
None => (),
Some(release) => unsafe { release(self) },
}
}
}
impl FFI_ArrowArrayStream {
pub fn new(batch_reader: Box<dyn RecordBatchReader + Send>) -> Self {
let private_data = Box::new(StreamPrivateData {
batch_reader,
last_error: None,
});
Self {
get_schema: Some(get_schema),
get_next: Some(get_next),
get_last_error: Some(get_last_error),
release: Some(release_stream),
private_data: Box::into_raw(private_data).cast::<c_void>(),
}
}
pub unsafe fn from_raw(raw_stream: *mut FFI_ArrowArrayStream) -> Self {
unsafe { std::ptr::replace(raw_stream, Self::empty()) }
}
pub fn empty() -> Self {
Self {
get_schema: None,
get_next: None,
get_last_error: None,
release: None,
private_data: std::ptr::null_mut(),
}
}
pub fn release(&self) -> Option<unsafe extern "C" fn(arg1: *mut Self)> {
self.release
}
pub fn private_data(&self) -> *mut c_void {
self.private_data
}
pub unsafe fn set_release(
&mut self,
release: Option<unsafe extern "C" fn(arg1: *mut Self)>,
) -> Option<unsafe extern "C" fn(arg1: *mut Self)> {
std::mem::replace(&mut self.release, release)
}
pub unsafe fn set_private_data(&mut self, private_data: *mut c_void) -> *mut c_void {
std::mem::replace(&mut self.private_data, private_data)
}
}
struct ExportedArrayStream {
stream: *mut FFI_ArrowArrayStream,
}
impl ExportedArrayStream {
fn get_private_data(&mut self) -> &mut StreamPrivateData {
unsafe { &mut *(*self.stream).private_data.cast::<StreamPrivateData>() }
}
pub fn get_schema(&mut self, out: *mut FFI_ArrowSchema) -> i32 {
let private_data = self.get_private_data();
let reader = &private_data.batch_reader;
let schema = FFI_ArrowSchema::try_from(reader.schema().as_ref());
match schema {
Ok(schema) => {
unsafe { std::ptr::copy(addr_of!(schema), out, 1) };
std::mem::forget(schema);
0
}
Err(ref err) => {
private_data.last_error = Some(
CString::new(err.to_string()).expect("Error string has a null byte in it."),
);
get_error_code(err)
}
}
}
pub fn get_next(&mut self, out: *mut FFI_ArrowArray) -> i32 {
let private_data = self.get_private_data();
let reader = &mut private_data.batch_reader;
match reader.next() {
None => {
unsafe { std::ptr::write(out, FFI_ArrowArray::empty()) }
0
}
Some(next_batch) => {
if let Ok(batch) = next_batch {
let struct_array = StructArray::from(batch);
let array = FFI_ArrowArray::new(&struct_array.to_data());
unsafe { std::ptr::write_unaligned(out, array) };
0
} else {
let err = &next_batch.unwrap_err();
private_data.last_error = Some(
CString::new(err.to_string()).expect("Error string has a null byte in it."),
);
get_error_code(err)
}
}
}
}
pub fn get_last_error(&mut self) -> Option<&CString> {
self.get_private_data().last_error.as_ref()
}
}
fn get_error_code(err: &ArrowError) -> i32 {
match err {
ArrowError::NotYetImplemented(_) => ENOSYS,
ArrowError::MemoryError(_) => ENOMEM,
ArrowError::IoError(_, _) => EIO,
_ => EINVAL,
}
}
#[derive(Debug)]
pub struct ArrowArrayStreamReader {
stream: FFI_ArrowArrayStream,
schema: SchemaRef,
}
unsafe fn producer_error(stream_ptr: *mut FFI_ArrowArrayStream) -> Option<String> {
let get_last_error = unsafe { (*stream_ptr).get_last_error }?;
let error_str = unsafe { get_last_error(stream_ptr) };
if error_str.is_null() {
return None;
}
Some(
unsafe { CStr::from_ptr(error_str) }
.to_string_lossy()
.into_owned(),
)
}
fn get_stream_schema(stream_ptr: *mut FFI_ArrowArrayStream) -> Result<SchemaRef> {
let mut schema = FFI_ArrowSchema::empty();
let ret_code = unsafe { (*stream_ptr).get_schema.unwrap()(stream_ptr, &raw mut schema) };
if ret_code == 0 {
let schema = Schema::try_from(&schema)?;
Ok(Arc::new(schema))
} else {
let message = format!("Cannot get schema from input stream. Error code: {ret_code}");
let message = match unsafe { producer_error(stream_ptr) } {
Some(producer_message) => format!("{message}. Producer error: {producer_message}"),
None => message,
};
Err(ArrowError::CDataInterface(message))
}
}
impl ArrowArrayStreamReader {
pub fn try_new(mut stream: FFI_ArrowArrayStream) -> Result<Self> {
if stream.release.is_none() {
return Err(ArrowError::CDataInterface(
"input stream is already released".to_string(),
));
}
let schema = get_stream_schema(&raw mut stream)?;
Ok(Self { stream, schema })
}
pub unsafe fn from_raw(raw_stream: *mut FFI_ArrowArrayStream) -> Result<Self> {
Self::try_new(unsafe { FFI_ArrowArrayStream::from_raw(raw_stream) })
}
}
impl Iterator for ArrowArrayStreamReader {
type Item = Result<RecordBatch>;
fn next(&mut self) -> Option<Self::Item> {
let mut array = FFI_ArrowArray::empty();
let ret_code =
unsafe { self.stream.get_next.unwrap()(&raw mut self.stream, &raw mut array) };
if ret_code == 0 {
if array.is_released() {
return None;
}
let result = unsafe {
from_ffi_and_data_type(array, DataType::Struct(self.schema().fields().clone()))
};
Some(result.and_then(|data| {
let len = data.len();
RecordBatch::try_new_with_options(
self.schema.clone(),
StructArray::from(data).into_parts().1,
&RecordBatchOptions::new().with_row_count(Some(len)),
)
}))
} else {
let message =
format!("Cannot get next batch from input stream. Error code: {ret_code}");
let message = match unsafe { producer_error(&raw mut self.stream) } {
Some(producer_message) => format!("{message}. Producer error: {producer_message}"),
None => message,
};
Some(Err(ArrowError::CDataInterface(message)))
}
}
}
impl RecordBatchReader for ArrowArrayStreamReader {
fn schema(&self) -> SchemaRef {
self.schema.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use arrow_schema::Field;
use crate::array::Int32Array;
use crate::ffi::from_ffi;
struct TestRecordBatchReader {
schema: SchemaRef,
iter: Box<dyn Iterator<Item = Result<RecordBatch>> + Send>,
}
impl TestRecordBatchReader {
pub fn new(
schema: SchemaRef,
iter: Box<dyn Iterator<Item = Result<RecordBatch>> + Send>,
) -> TestRecordBatchReader {
TestRecordBatchReader { schema, iter }
}
}
impl Iterator for TestRecordBatchReader {
type Item = Result<RecordBatch>;
fn next(&mut self) -> Option<Self::Item> {
self.iter.next()
}
}
impl RecordBatchReader for TestRecordBatchReader {
fn schema(&self) -> SchemaRef {
self.schema.clone()
}
}
fn _test_round_trip_export(batch: RecordBatch, schema: Arc<Schema>) -> Result<()> {
let iter = Box::new(vec![batch.clone(), batch.clone()].into_iter().map(Ok)) as _;
let reader = Box::new(TestRecordBatchReader::new(schema.clone(), iter));
let mut ffi_stream = FFI_ArrowArrayStream::new(reader);
let mut ffi_schema = FFI_ArrowSchema::empty();
let ret_code = unsafe { get_schema(&raw mut ffi_stream, &raw mut ffi_schema) };
assert_eq!(ret_code, 0);
let exported_schema = Schema::try_from(&ffi_schema).unwrap();
assert_eq!(&exported_schema, schema.as_ref());
let mut produced_batches = vec![];
loop {
let mut ffi_array = FFI_ArrowArray::empty();
let ret_code = unsafe { get_next(&raw mut ffi_stream, &raw mut ffi_array) };
assert_eq!(ret_code, 0);
if ffi_array.is_released() {
break;
}
let array = unsafe { from_ffi(ffi_array, &ffi_schema) }.unwrap();
let len = array.len();
let record_batch = RecordBatch::try_new_with_options(
SchemaRef::from(exported_schema.clone()),
StructArray::from(array).into_parts().1,
&RecordBatchOptions::new().with_row_count(Some(len)),
)
.unwrap();
produced_batches.push(record_batch);
}
assert_eq!(produced_batches, vec![batch.clone(), batch]);
Ok(())
}
fn _test_round_trip_import(batch: RecordBatch, schema: Arc<Schema>) -> Result<()> {
let iter = Box::new(vec![batch.clone(), batch.clone()].into_iter().map(Ok)) as _;
let reader = Box::new(TestRecordBatchReader::new(schema.clone(), iter));
let stream = FFI_ArrowArrayStream::new(reader);
let stream_reader = ArrowArrayStreamReader::try_new(stream).unwrap();
let imported_schema = stream_reader.schema();
assert_eq!(imported_schema, schema);
let mut produced_batches = vec![];
for batch in stream_reader {
produced_batches.push(batch.unwrap());
}
assert_eq!(produced_batches, vec![batch.clone(), batch]);
Ok(())
}
#[test]
fn test_stream_round_trip() {
let array = Int32Array::from(vec![Some(2), None, Some(1), None]);
let array: Arc<dyn Array> = Arc::new(array);
let metadata = HashMap::from([("foo".to_owned(), "bar".to_owned())]);
let schema = Arc::new(Schema::new_with_metadata(
vec![
Field::new("a", array.data_type().clone(), true).with_metadata(metadata.clone()),
Field::new("b", array.data_type().clone(), true).with_metadata(metadata.clone()),
Field::new("c", array.data_type().clone(), true).with_metadata(metadata.clone()),
],
metadata,
));
let batch = RecordBatch::try_new(schema.clone(), vec![array.clone(), array.clone(), array])
.unwrap();
_test_round_trip_export(batch.clone(), schema.clone()).unwrap();
_test_round_trip_import(batch, schema).unwrap();
}
#[test]
fn test_stream_round_trip_no_columns() {
let metadata = HashMap::from([("foo".to_owned(), "bar".to_owned())]);
let schema = Arc::new(Schema::new_with_metadata(Vec::<Field>::new(), metadata));
let batch = RecordBatch::try_new_with_options(
schema.clone(),
Vec::<Arc<dyn Array>>::new(),
&RecordBatchOptions::new().with_row_count(Some(10)),
)
.unwrap();
_test_round_trip_export(batch.clone(), schema.clone()).unwrap();
_test_round_trip_import(batch, schema).unwrap();
}
#[test]
fn test_error_import() -> Result<()> {
let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)]));
let iter =
Box::new(vec![Err(ArrowError::MemoryError("out of memory".to_string()))].into_iter());
let reader = Box::new(TestRecordBatchReader::new(schema.clone(), iter));
let stream = FFI_ArrowArrayStream::new(reader);
let stream_reader = ArrowArrayStreamReader::try_new(stream).unwrap();
let imported_schema = stream_reader.schema();
assert_eq!(imported_schema, schema);
let mut produced_batches = vec![];
for batch in stream_reader {
produced_batches.push(batch);
}
assert_eq!(produced_batches.len(), 1);
assert_eq!(
produced_batches[0].as_ref().unwrap_err().to_string(),
format!(
"C Data interface error: Cannot get next batch from input stream. \
Error code: {ENOMEM}. Producer error: Memory error: out of memory"
)
);
Ok(())
}
unsafe extern "C" fn failing_get_schema(
_stream: *mut FFI_ArrowArrayStream,
_out: *mut FFI_ArrowSchema,
) -> c_int {
EIO
}
unsafe extern "C" fn working_get_schema(
_stream: *mut FFI_ArrowArrayStream,
out: *mut FFI_ArrowSchema,
) -> c_int {
let schema = Schema::new(vec![Field::new("a", DataType::Int32, true)]);
unsafe { std::ptr::write(out, FFI_ArrowSchema::try_from(&schema).unwrap()) };
0
}
unsafe extern "C" fn failing_get_next(
_stream: *mut FFI_ArrowArrayStream,
_out: *mut FFI_ArrowArray,
) -> c_int {
EIO
}
unsafe extern "C" fn producer_last_error(_stream: *mut FFI_ArrowArrayStream) -> *const c_char {
c"the producer failed".as_ptr()
}
unsafe extern "C" fn null_last_error(_stream: *mut FFI_ArrowArrayStream) -> *const c_char {
std::ptr::null()
}
unsafe extern "C" fn mark_released(stream: *mut FFI_ArrowArrayStream) {
unsafe { (*stream).release = None };
}
fn failing_stream(
get_last_error: Option<unsafe extern "C" fn(*mut FFI_ArrowArrayStream) -> *const c_char>,
) -> FFI_ArrowArrayStream {
let mut stream = FFI_ArrowArrayStream::empty();
stream.get_schema = Some(failing_get_schema);
stream.get_next = Some(failing_get_next);
stream.get_last_error = get_last_error;
stream.release = Some(mark_released);
stream
}
#[test]
fn test_import_schema_error_reports_producer_message() {
let err =
ArrowArrayStreamReader::try_new(failing_stream(Some(producer_last_error))).unwrap_err();
assert_eq!(
err.to_string(),
format!(
"C Data interface error: Cannot get schema from input stream. \
Error code: {EIO}. Producer error: the producer failed"
)
);
}
#[test]
fn test_import_schema_error_without_producer_message() {
let err =
ArrowArrayStreamReader::try_new(failing_stream(Some(null_last_error))).unwrap_err();
assert_eq!(
err.to_string(),
format!(
"C Data interface error: Cannot get schema from input stream. Error code: {EIO}"
)
);
}
#[test]
fn test_import_schema_error_without_error_callback() {
let err = ArrowArrayStreamReader::try_new(failing_stream(None)).unwrap_err();
assert_eq!(
err.to_string(),
format!(
"C Data interface error: Cannot get schema from input stream. Error code: {EIO}"
)
);
}
#[test]
fn test_import_next_error_without_producer_message() {
let mut stream = failing_stream(Some(null_last_error));
stream.get_schema = Some(working_get_schema);
let err = ArrowArrayStreamReader::try_new(stream)
.unwrap()
.next()
.unwrap()
.unwrap_err();
assert_eq!(
err.to_string(),
format!(
"C Data interface error: Cannot get next batch from input stream. Error code: {EIO}"
)
);
}
static STREAM_WRAPPER_RAN: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
struct StreamWrapperData {
original_release: Option<unsafe extern "C" fn(*mut FFI_ArrowArrayStream)>,
original_private_data: *mut c_void,
}
unsafe extern "C" fn wrapping_release(stream: *mut FFI_ArrowArrayStream) {
use std::sync::atomic::Ordering;
let stream = unsafe { &mut *stream };
let data = unsafe { Box::from_raw(stream.private_data().cast::<StreamWrapperData>()) };
STREAM_WRAPPER_RAN.store(true, Ordering::SeqCst);
unsafe { stream.set_release(data.original_release) };
unsafe { stream.set_private_data(data.original_private_data) };
if let Some(release) = stream.release() {
unsafe { release(stream) };
}
}
#[test]
fn test_wrap_release_callback() {
use std::sync::atomic::Ordering;
let batch_reader = Box::new(TestRecordBatchReader::new(
Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)])),
Box::new(std::iter::empty()),
));
let mut stream = FFI_ArrowArrayStream::new(batch_reader);
let data = Box::new(StreamWrapperData {
original_release: stream.release(),
original_private_data: stream.private_data(),
});
unsafe { stream.set_release(Some(wrapping_release)) };
unsafe { stream.set_private_data(Box::into_raw(data).cast::<c_void>()) };
drop(stream); assert!(STREAM_WRAPPER_RAN.load(Ordering::SeqCst));
}
}