use crate::{
DuckBindArgs, DuckDynamicRow, DuckExtraInfo, DuckResult, DuckResultSchema, DuckTypeDesc,
duck_error, panic_to_string, raw_extra_info,
};
use libduckdb_sys::{
DuckDBSuccess, duckdb_bind_info, duckdb_copy_function, duckdb_copy_function_set_copy_from_function,
duckdb_copy_function_set_name, duckdb_create_copy_function, duckdb_create_table_function,
duckdb_data_chunk, duckdb_data_chunk_set_size, duckdb_destroy_copy_function,
duckdb_destroy_table_function, duckdb_function_info, duckdb_init_info,
duckdb_init_set_init_data, duckdb_register_copy_function, duckdb_table_function,
duckdb_table_function_add_named_parameter, duckdb_table_function_add_parameter,
duckdb_table_function_bind_get_result_column_count,
duckdb_table_function_bind_get_result_column_name,
duckdb_table_function_bind_get_result_column_type, duckdb_table_function_set_bind,
duckdb_table_function_set_extra_info, duckdb_table_function_set_function,
duckdb_table_function_set_init, duckdb_table_function_set_name,
};
use quack_rs::connection::Connection;
use quack_rs::data_chunk::DataChunk;
use quack_rs::prelude::{LogicalType, TypeId};
use quack_rs::table::{BindInfo, FfiBindData, FfiInitData, FunctionInfo, InitInfo};
use quack_rs::vector::vector_size;
use std::ffi::{CStr, CString};
use std::os::raw::c_void;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::sync::Mutex;
pub trait DuckCopyFromReader: Sized + Send + 'static {
type Args: DuckBindArgs;
fn open(args: Self::Args, schema: &DuckResultSchema) -> DuckResult<Self>;
fn finish(&mut self) -> DuckResult<()> {
Ok(())
}
}
pub trait CopyFromFunctionAdapter: Sized + 'static {
const NAME: &'static str;
type Reader: DuckCopyFromReader;
fn extra_info() -> Option<DuckExtraInfo> {
None
}
fn next_batch(reader: &mut Self::Reader, limit: usize) -> DuckResult<Vec<DuckDynamicRow>>;
unsafe extern "C" fn c_bind(raw: duckdb_bind_info) {
let info = unsafe { BindInfo::new(raw) };
handle(
AssertUnwindSafe(|| {
ensure_single_path_parameter::<Self>()?;
let args = <Self::Reader as DuckCopyFromReader>::Args::read_bind_args(&info)?;
let schema = unsafe { read_target_schema(raw)? };
let reader = Self::Reader::open(args, &schema)?;
let state = CopyFromState {
schema,
reader,
};
unsafe {
FfiBindData::<Mutex<Option<CopyFromState<Self::Reader>>>>::set(
raw,
Mutex::new(Some(state)),
);
}
Ok(())
}),
|message| info.set_error(message),
);
}
unsafe extern "C" fn c_init(raw: duckdb_init_info) {
let info = unsafe { InitInfo::new(raw) };
handle(
AssertUnwindSafe(|| {
let cell = unsafe {
FfiBindData::<Mutex<Option<CopyFromState<Self::Reader>>>>::get_from_init(raw)
};
let Some(cell) = cell else {
return Err(duck_error(format!(
"{}: missing reader state (was the bind callback skipped?)",
Self::NAME
)));
};
let mut guard = cell.lock().map_err(|_| {
duck_error(format!("{}: reader state mutex was poisoned", Self::NAME))
})?;
let Some(state) = guard.take() else {
return Err(duck_error(format!(
"{}: reader state was already consumed",
Self::NAME
)));
};
unsafe {
duckdb_init_set_init_data(
raw,
Box::into_raw(Box::new(state)).cast::<c_void>(),
Some(destroy_copy_from_state::<Self>),
);
}
info.set_max_threads(1);
Ok(())
}),
|message| info.set_error(message),
);
}
unsafe extern "C" fn c_scan(raw: duckdb_function_info, output: duckdb_data_chunk) {
let info = unsafe { FunctionInfo::new(raw) };
let result = catch_unwind(AssertUnwindSafe(|| {
let state = unsafe { FfiInitData::<CopyFromState<Self::Reader>>::get_mut(raw) };
let Some(state) = state else {
return Err(duck_error(format!(
"{}: missing reader state (was the init callback skipped?)",
Self::NAME
)));
};
let rows = Self::next_batch(&mut state.reader, vector_size() as usize)?;
if rows.is_empty() {
unsafe { duckdb_data_chunk_set_size(output, 0) };
return Ok(());
}
let chunk = unsafe { DataChunk::from_raw(output) };
let refs: Vec<Option<&DuckDynamicRow>> = rows.iter().map(Some).collect();
DuckDynamicRow::write_batch(&chunk, &state.schema, &refs)?;
unsafe { chunk.set_size(rows.len()) };
Ok(())
}));
match result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
info.set_error(e.as_str());
unsafe { duckdb_data_chunk_set_size(output, 0) };
}
Err(payload) => {
info.set_error(&panic_to_string(payload));
unsafe { duckdb_data_chunk_set_size(output, 0) };
}
}
}
unsafe fn reader_table_function() -> DuckResult<duckdb_table_function> {
let name = CString::new(Self::NAME)
.map_err(|_| duck_error(format!("{}: format name contains a NUL byte", Self::NAME)))?;
let mut named = Vec::new();
for (param_name, logical_type) in
<Self::Reader as DuckCopyFromReader>::Args::bind_param_logical()
{
if let Some(param_name) = param_name {
let c_name = CString::new(param_name.as_str()).map_err(|_| {
duck_error(format!(
"{}: COPY option name contains a NUL byte",
Self::NAME
))
})?;
named.push((c_name, logical_type));
}
}
let function = unsafe { duckdb_create_table_function() };
if function.is_null() {
return Err(duck_error(format!(
"{}: duckdb_create_table_function returned null",
Self::NAME
)));
}
unsafe {
duckdb_table_function_set_name(function, name.as_ptr());
let path_type = LogicalType::new(TypeId::Varchar);
duckdb_table_function_add_parameter(function, path_type.as_raw());
for (c_name, logical_type) in &named {
duckdb_table_function_add_named_parameter(
function,
c_name.as_ptr(),
logical_type.as_raw(),
);
}
duckdb_table_function_set_bind(function, Some(Self::c_bind));
duckdb_table_function_set_init(function, Some(Self::c_init));
duckdb_table_function_set_function(function, Some(Self::c_scan));
if let Some((ptr, destroy)) = raw_extra_info(Self::extra_info()) {
duckdb_table_function_set_extra_info(function, ptr, destroy);
}
}
Ok(function)
}
unsafe fn register(c: &Connection) -> DuckResult<()> {
let reader = unsafe { Self::reader_table_function()? };
let copy_function = unsafe { duckdb_create_copy_function() };
if copy_function.is_null() {
let mut reader = reader;
unsafe { duckdb_destroy_table_function(&raw mut reader) };
return Err(duck_error(format!(
"{}: duckdb_create_copy_function returned null",
Self::NAME
)));
}
let name = CString::new(Self::NAME)
.map_err(|_| duck_error(format!("{}: format name contains a NUL byte", Self::NAME)))?;
unsafe {
duckdb_copy_function_set_name(copy_function, name.as_ptr());
duckdb_copy_function_set_copy_from_function(copy_function, reader);
}
let result = unsafe { duckdb_register_copy_function(c.as_raw_connection(), copy_function) };
let mut copy_function: duckdb_copy_function = copy_function;
unsafe { duckdb_destroy_copy_function(&raw mut copy_function) };
if result == DuckDBSuccess {
Ok(())
} else {
Err(duck_error(format!(
"{}: duckdb_register_copy_function failed",
Self::NAME
)))
}
}
}
struct CopyFromState<R: DuckCopyFromReader> {
schema: DuckResultSchema,
reader: R,
}
fn ensure_single_path_parameter<A: CopyFromFunctionAdapter>() -> DuckResult<()> {
let positional =
<<A as CopyFromFunctionAdapter>::Reader as DuckCopyFromReader>::Args::bind_param_logical()
.iter()
.filter(|(name, _)| name.is_none())
.count();
if positional == 1 {
return Ok(());
}
Err(duck_error(format!(
"{}: a COPY FROM reader must declare exactly one positional parameter (the file path), but \
its arguments declare {positional}",
A::NAME
)))
}
unsafe fn read_target_schema(raw: duckdb_bind_info) -> DuckResult<DuckResultSchema> {
let count = unsafe { duckdb_table_function_bind_get_result_column_count(raw) };
if count == 0 {
return Err(duck_error(
"copy from: the target table has no columns; a COPY FROM reader reads the target \
schema instead of declaring result columns",
));
}
let mut columns = Vec::with_capacity(count as usize);
for index in 0..count {
let name_ptr = unsafe { duckdb_table_function_bind_get_result_column_name(raw, index) };
let name = if name_ptr.is_null() {
format!("column_{index}")
} else {
unsafe { CStr::from_ptr(name_ptr) }
.to_string_lossy()
.into_owned()
};
let logical_type =
unsafe { LogicalType::from_raw(duckdb_table_function_bind_get_result_column_type(raw, index)) };
columns.push((name, DuckTypeDesc::from_logical_type(&logical_type)?));
}
Ok(DuckResultSchema::new(columns))
}
unsafe extern "C" fn destroy_copy_from_state<A: CopyFromFunctionAdapter>(ptr: *mut c_void) {
if ptr.is_null() {
return;
}
let mut state = unsafe { Box::from_raw(ptr.cast::<CopyFromState<A::Reader>>()) };
if let Err(e) = state.reader.finish() {
eprintln!(
"-- [duckfn] {}: reader finish failed at end of query: {}",
A::NAME,
e.as_str()
);
}
drop(state);
}
fn handle<F>(f: AssertUnwindSafe<F>, set_error: impl FnOnce(&str))
where
F: FnOnce() -> DuckResult<()>,
{
match catch_unwind(f) {
Ok(Ok(())) => {}
Ok(Err(e)) => set_error(e.as_str()),
Err(payload) => set_error(&panic_to_string(payload)),
}
}