use crate::{
DuckDynamicRow, DuckDynamicValue, DuckResult, DuckResultSchema, DuckTypeDesc, duck_error,
duck_value_is_null, panic_to_string,
};
use libduckdb_sys::{
duckdb_copy_function_bind_get_options, duckdb_copy_function_bind_info,
duckdb_copy_function_finalize_info, duckdb_copy_function_global_init_info,
duckdb_copy_function_sink_info, duckdb_data_chunk, duckdb_get_value_type,
};
use quack_rs::connection::Connection;
use quack_rs::copy_function::{
CopyBindInfo, CopyFinalizeInfo, CopyFunctionBuilder, CopyGlobalInitInfo, CopySinkInfo,
};
use quack_rs::data_chunk::DataChunk;
use quack_rs::prelude::{LogicalType, Value};
use std::os::raw::c_void;
use std::panic::{AssertUnwindSafe, catch_unwind};
#[derive(Debug, Clone, Default, PartialEq)]
pub struct DuckCopyOptions {
entries: Vec<(String, Option<DuckDynamicValue>)>,
}
impl DuckCopyOptions {
#[must_use]
pub fn new(entries: Vec<(String, Option<DuckDynamicValue>)>) -> Self {
Self { entries }
}
#[must_use]
pub fn entries(&self) -> &[(String, Option<DuckDynamicValue>)] {
&self.entries
}
#[must_use]
pub fn get(&self, name: &str) -> Option<&DuckDynamicValue> {
self.entries
.iter()
.find(|(key, _)| key.eq_ignore_ascii_case(name))
.and_then(|(_, value)| value.as_ref())
}
#[must_use]
pub fn contains(&self, name: &str) -> bool {
self.entries
.iter()
.any(|(key, _)| key.eq_ignore_ascii_case(name))
}
#[must_use]
pub fn get_bool(&self, name: &str) -> Option<bool> {
match self.get(name) {
Some(DuckDynamicValue::Boolean(value)) => Some(*value),
_ => None,
}
}
#[must_use]
pub fn get_i64(&self, name: &str) -> Option<i64> {
match self.get(name) {
Some(DuckDynamicValue::TinyInt(value)) => Some(i64::from(*value)),
Some(DuckDynamicValue::SmallInt(value)) => Some(i64::from(*value)),
Some(DuckDynamicValue::Integer(value)) => Some(i64::from(*value)),
Some(DuckDynamicValue::BigInt(value)) => Some(*value),
Some(DuckDynamicValue::UTinyInt(value)) => Some(i64::from(*value)),
Some(DuckDynamicValue::USmallInt(value)) => Some(i64::from(*value)),
Some(DuckDynamicValue::UInteger(value)) => Some(i64::from(*value)),
Some(DuckDynamicValue::UBigInt(value)) => i64::try_from(*value).ok(),
_ => None,
}
}
#[must_use]
pub fn get_str(&self, name: &str) -> Option<&str> {
match self.get(name) {
Some(DuckDynamicValue::Varchar(value)) => Some(value.as_str()),
_ => None,
}
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
}
pub trait DuckCopyToWriter: Sized + 'static {
fn open(path: &str, schema: &DuckResultSchema, options: &DuckCopyOptions) -> DuckResult<Self>;
fn finish(&mut self) -> DuckResult<()> {
Ok(())
}
}
pub trait CopyToFunctionAdapter: Sized + 'static {
const NAME: &'static str;
type Writer: DuckCopyToWriter;
fn write_rows(writer: &mut Self::Writer, rows: &[DuckDynamicRow]) -> DuckResult<()>;
unsafe extern "C" fn c_bind(raw: duckdb_copy_function_bind_info) {
let info = unsafe { CopyBindInfo::new(raw) };
handle(
AssertUnwindSafe(|| {
let schema = Self::read_schema(&info)?;
let options = unsafe { read_copy_options(raw)? };
let boxed = Box::new(CopyToBindData { schema, options });
unsafe {
info.set_bind_data(
Box::into_raw(boxed).cast::<c_void>(),
Some(drop_boxed::<CopyToBindData>),
);
}
Ok(())
}),
|message| info.set_error(message),
);
}
unsafe extern "C" fn c_global_init(info: duckdb_copy_function_global_init_info) {
let info = unsafe { CopyGlobalInitInfo::new(info) };
handle(
AssertUnwindSafe(|| {
let path = unsafe { info.get_file_path() };
let bind = bind_data_of(&info)?;
let writer = Self::Writer::open(&path, &bind.schema, &bind.options)?;
let boxed = Box::new(writer);
unsafe {
info.set_global_state(
Box::into_raw(boxed).cast::<c_void>(),
Some(drop_boxed::<Self::Writer>),
);
}
Ok(())
}),
|message| info.set_error(message),
);
}
unsafe extern "C" fn c_sink(raw: duckdb_copy_function_sink_info, raw_chunk: duckdb_data_chunk) {
let info = unsafe { CopySinkInfo::new(raw) };
handle(
AssertUnwindSafe(|| {
let writer = writer_of::<Self::Writer>(&info)?;
let bind = bind_data_of(&info)?;
let chunk = unsafe { DataChunk::from_raw(raw_chunk) };
let rows = DuckDynamicRow::read_batch(&chunk, &bind.schema)?;
Self::write_rows(writer, &rows)
}),
|message| info.set_error(message),
);
}
unsafe extern "C" fn c_finalize(info: duckdb_copy_function_finalize_info) {
let info = unsafe { CopyFinalizeInfo::new(info) };
handle(
AssertUnwindSafe(|| {
let state_ptr = unsafe { info.get_global_state() }.cast::<Self::Writer>();
if state_ptr.is_null() {
return Ok(());
}
let writer = unsafe { &mut *state_ptr };
writer.finish()
}),
|message| info.set_error(message),
);
}
fn read_schema(info: &CopyBindInfo) -> DuckResult<DuckResultSchema> {
let count = info.column_count();
let mut columns = Vec::with_capacity(count as usize);
for index in 0..count {
let logical_type = unsafe { info.column_type(index) };
columns.push((
format!("column_{index}"),
DuckTypeDesc::from_logical_type(&logical_type)?,
));
}
Ok(DuckResultSchema::new(columns))
}
fn copy_function_builder() -> DuckResult<CopyFunctionBuilder> {
let builder = CopyFunctionBuilder::try_new(Self::NAME)?
.bind(Self::c_bind)
.global_init(Self::c_global_init)
.sink(Self::c_sink)
.finalize(Self::c_finalize);
Ok(builder)
}
unsafe fn register(c: &Connection) -> DuckResult<()> {
let builder = Self::copy_function_builder()?;
unsafe { builder.register(c.as_raw_connection()) }
}
}
struct CopyToBindData {
schema: DuckResultSchema,
options: DuckCopyOptions,
}
unsafe fn read_copy_options(raw: duckdb_copy_function_bind_info) -> DuckResult<DuckCopyOptions> {
let value = unsafe { Value::from_raw(duckdb_copy_function_bind_get_options(raw)) };
let options = parse_copy_options(&value);
let _ = value.into_raw();
options
}
fn parse_copy_options(value: &Value) -> DuckResult<DuckCopyOptions> {
if duck_value_is_null(value) {
return Ok(DuckCopyOptions::default());
}
let logical_type = unsafe { LogicalType::from_raw(duckdb_get_value_type(value.as_raw())) };
let count = unsafe { logical_type.struct_child_count() };
let mut entries = Vec::with_capacity(count as usize);
for index in 0..count {
let name = unsafe { logical_type.struct_child_name(index) };
let field_type = unsafe { logical_type.struct_child_type(index) };
let Ok(desc) = DuckTypeDesc::from_logical_type(&field_type) else {
continue;
};
let child = match value.struct_child(index as usize) {
Some(child) => DuckDynamicValue::from_duck_value(&child, &desc)?,
None => None,
};
entries.push((name, child));
}
Ok(DuckCopyOptions::new(entries))
}
fn bind_data_of(info: &impl BindDataAccessor) -> DuckResult<&CopyToBindData> {
let ptr = unsafe { info.bind_data() }.cast::<CopyToBindData>();
if ptr.is_null() {
return Err(duck_error(
"copy function: bind data was not initialized; was the bind callback skipped?",
));
}
Ok(unsafe { &*ptr })
}
fn writer_of<W>(info: &impl GlobalStateAccessor) -> DuckResult<&mut W> {
let ptr = unsafe { info.global_state() }.cast::<W>();
if ptr.is_null() {
return Err(duck_error(
"copy function: writer was not initialized; was global init skipped?",
));
}
Ok(unsafe { &mut *ptr })
}
trait BindDataAccessor {
unsafe fn bind_data(&self) -> *mut c_void;
}
impl BindDataAccessor for CopyGlobalInitInfo {
unsafe fn bind_data(&self) -> *mut c_void {
unsafe { self.get_bind_data() }
}
}
impl BindDataAccessor for CopySinkInfo {
unsafe fn bind_data(&self) -> *mut c_void {
unsafe { self.get_bind_data() }
}
}
trait GlobalStateAccessor {
unsafe fn global_state(&self) -> *mut c_void;
}
impl GlobalStateAccessor for CopySinkInfo {
unsafe fn global_state(&self) -> *mut c_void {
unsafe { self.get_global_state() }
}
}
impl GlobalStateAccessor for CopyFinalizeInfo {
unsafe fn global_state(&self) -> *mut c_void {
unsafe { self.get_global_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)),
}
}
unsafe extern "C" fn drop_boxed<T>(ptr: *mut c_void) {
if !ptr.is_null() {
drop(unsafe { Box::from_raw(ptr.cast::<T>()) });
}
}