use crate::value_types::duck_value_type::DuckValueType;
use crate::{
DuckExtraInfo, DuckOptionResult, DuckResult, erased_extra_info, panic_to_string, raw_extra_info,
vec_option_to_ref,
};
use libduckdb_sys::{duckdb_function_info, duckdb_vector, idx_t};
use quack_rs::prelude::{CastFunctionBuilder, CastFunctionInfo, CastMode, Connection, Registrar};
use std::panic::{AssertUnwindSafe, catch_unwind};
pub trait CastFunctionAdapter: Sized + 'static {
const NAME: &'static str;
type Input: DuckValueType;
type Output: DuckValueType;
fn implicit_cost() -> Option<i64> {
None
}
fn extra_info() -> Option<DuckExtraInfo> {
None
}
fn apply_with_null(value: Option<Self::Input>) -> DuckOptionResult<Self::Output>;
fn apply_with_extra(
value: Option<Self::Input>,
extra: Option<&DuckExtraInfo>,
) -> DuckOptionResult<Self::Output> {
let _ = extra;
Self::apply_with_null(value)
}
fn cast_function_builder() -> CastFunctionBuilder {
let mut builder = CastFunctionBuilder::new_logical(
Self::Input::logical_type(),
Self::Output::logical_type(),
)
.function(Self::cast_function_wrapper);
if let Some(cost) = Self::implicit_cost() {
builder = builder.implicit_cost(cost);
}
if let Some((ptr, destroy)) = raw_extra_info(Self::extra_info()) {
builder = unsafe { builder.extra_info(ptr, destroy) };
}
builder
}
unsafe fn register(c: &Connection) -> DuckResult<()> {
unsafe { c.register_cast(Self::cast_function_builder()) }
}
unsafe extern "C" fn cast_function_wrapper(
info: duckdb_function_info,
count: idx_t,
input: duckdb_vector,
output: duckdb_vector,
) -> bool {
let info = unsafe { CastFunctionInfo::new(info) };
match catch_unwind(AssertUnwindSafe(|| {
cast_batch::<Self>(&info, count as usize, input, output)
})) {
Ok(success) => success,
Err(e) => {
info.set_error(&panic_to_string(e));
false
}
}
}
}
fn cast_batch<A: CastFunctionAdapter>(
info: &CastFunctionInfo,
count: usize,
input: duckdb_vector,
output: duckdb_vector,
) -> bool {
let reader = A::Input::create_reader_from_vector(input, count);
let try_mode = info.cast_mode() == CastMode::Try;
let extra = unsafe { erased_extra_info(info) };
let mut results: Vec<Option<A::Output>> = Vec::with_capacity(count);
for row in 0..count {
let value = A::Input::read(&reader, row);
match A::apply_with_extra(value, extra) {
Ok(v) => results.push(v),
Err(e) => {
if try_mode {
unsafe { info.set_row_error(e.as_str(), row as idx_t, output) };
results.push(None);
} else {
info.set_error(e.as_str());
return false;
}
}
}
}
A::Output::write_batch(output, &vec_option_to_ref(&results));
true
}