#[cfg(test)]
mod tests;
use super::{BindInfo, DataChunkHandle, InitInfo, LogicalTypeHandle, TableFunctionInfo, VTab};
use std::sync::{Arc, Mutex, OnceLock, atomic::AtomicUsize};
use arrow::{
array::{ArrayData, StructArray},
ffi::{FFI_ArrowArray, FFI_ArrowSchema, from_ffi},
record_batch::RecordBatch,
};
pub use crate::arrow_interop::*;
use crate::core::LogicalTypeId;
#[repr(C)]
pub struct ArrowBindData {
rb: RecordBatch,
}
#[repr(C)]
pub struct ArrowInitData {
offset: AtomicUsize,
vector_size: usize,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
struct RowSlice {
offset: usize,
len: usize,
}
impl ArrowInitData {
fn take_slice(&self, num_rows: usize) -> Option<RowSlice> {
use std::sync::atomic::Ordering::Relaxed;
let vector_size = self.vector_size;
let offset = self
.offset
.fetch_update(Relaxed, Relaxed, |offset| {
if offset >= num_rows {
None
} else {
Some(offset.saturating_add(vector_size).min(num_rows))
}
})
.ok()?;
Some(RowSlice {
offset,
len: (num_rows - offset).min(vector_size),
})
}
}
pub struct ArrowVTab;
const ARROW_QUERY_PARAMS_MARKER: usize = 0x4152_5257; fn arrow_record_batch_store() -> &'static Mutex<Vec<Arc<RecordBatch>>> {
static STORE: OnceLock<Mutex<Vec<Arc<RecordBatch>>>> = OnceLock::new();
STORE.get_or_init(|| Mutex::new(Vec::new()))
}
fn register_arrow_record_batch(rb: RecordBatch) -> [usize; 2] {
let mut store = arrow_record_batch_store()
.lock()
.expect("ArrowVTab record batch store poisoned");
let rb = Arc::new(rb);
let ptr = Arc::as_ptr(&rb);
store.push(rb);
[ptr as usize, ARROW_QUERY_PARAMS_MARKER]
}
fn arrow_query_param_usize(bind: &BindInfo, index: u64, name: &str) -> Result<usize, Box<dyn std::error::Error>> {
let value = bind.get_parameter(index);
if value.is_null() {
return Err(format!("ArrowVTab {name} parameter must not be NULL").into());
}
let logical_type = value.logical_type_id();
if logical_type != LogicalTypeId::UBigint {
return Err(format!("ArrowVTab {name} parameter must be UBIGINT, got {logical_type:?}").into());
}
usize::try_from(value.to_uint64()).map_err(|_| format!("ArrowVTab {name} parameter does not fit in usize").into())
}
unsafe fn address_to_arrow_record_batch(
address: usize,
marker: usize,
) -> Result<RecordBatch, Box<dyn std::error::Error>> {
let ptr = address as *const RecordBatch;
if ptr.is_null() {
return Err("invalid ArrowVTab record batch address".into());
}
if marker != ARROW_QUERY_PARAMS_MARKER {
return Err("ArrowVTab query parameter marker mismatch; use arrow_recordbatch_to_query_params".into());
}
Ok(unsafe { (*ptr).clone() })
}
fn arrow_record_batch_from_ffi(array: FFI_ArrowArray, schema: FFI_ArrowSchema) -> RecordBatch {
let array_data = unsafe { from_ffi(array, &schema) }.expect("failed to import Arrow FFI data");
let struct_array = StructArray::from(array_data);
RecordBatch::from(&struct_array)
}
impl VTab for ArrowVTab {
type BindData = ArrowBindData;
type InitData = ArrowInitData;
fn bind(bind: &BindInfo) -> Result<Self::BindData, Box<dyn std::error::Error>> {
let param_count = bind.get_parameter_count();
if param_count != 2 {
return Err(format!("Bad param count: {param_count}, expected 2").into());
}
let address = arrow_query_param_usize(bind, 0, "record batch address")?;
let marker = arrow_query_param_usize(bind, 1, "marker")?;
let rb = unsafe { address_to_arrow_record_batch(address, marker)? };
for f in rb.schema().fields() {
let name = f.name();
let logical_type = to_duckdb_logical_type_for_field(f)?;
bind.add_result_column(name, logical_type);
}
Ok(ArrowBindData { rb })
}
fn init(_: &InitInfo) -> Result<Self::InitData, Box<dyn std::error::Error>> {
let vector_size = unsafe { crate::ffi::duckdb_vector_size() } as usize;
if vector_size == 0 {
return Err("DuckDB vector size must be greater than zero".into());
}
Ok(ArrowInitData {
offset: AtomicUsize::new(0),
vector_size,
})
}
fn func(func: &TableFunctionInfo<Self>, output: &mut DataChunkHandle) -> Result<(), Box<dyn std::error::Error>> {
let init_info = func.get_init_data();
let bind_info = func.get_bind_data();
let rb = &bind_info.rb;
let num_rows = rb.num_rows();
let Some(slice) = init_info.take_slice(num_rows) else {
output.set_len(0);
return Ok(());
};
record_batch_to_duckdb_data_chunk(&rb.slice(slice.offset, slice.len), output)?;
Ok(())
}
fn parameters() -> Option<Vec<LogicalTypeHandle>> {
Some(vec![
LogicalTypeHandle::from(LogicalTypeId::UBigint), LogicalTypeHandle::from(LogicalTypeId::UBigint), ])
}
}
pub fn arrow_recordbatch_to_query_params(rb: RecordBatch) -> [usize; 2] {
register_arrow_record_batch(rb)
}
pub fn arrow_arraydata_to_query_params(data: ArrayData) -> [usize; 2] {
let struct_array = StructArray::from(data);
arrow_recordbatch_to_query_params(RecordBatch::from(&struct_array))
}
pub fn arrow_ffi_to_query_params(array: FFI_ArrowArray, schema: FFI_ArrowSchema) -> [usize; 2] {
arrow_recordbatch_to_query_params(arrow_record_batch_from_ffi(array, schema))
}