use crate::duck_columns::DuckColumns;
use crate::utils::builder_with_params::BuilderWithParams;
use crate::value_types::duck_value_type::{DuckValueReader, DuckValueType};
use crate::{
DuckExtraInfo, DuckOptionResult, DuckResult, duck_error, duck_scalar_unwind, erased_extra_info,
raw_extra_info, vec_option_to_ref,
};
use libduckdb_sys::{duckdb_connection, duckdb_data_chunk, duckdb_function_info, duckdb_vector};
use quack_rs::data_chunk::DataChunk;
use quack_rs::prelude::{
LogicalType, NullHandling, ScalarFunctionBuilder, ScalarFunctionInfo, ScalarOverloadBuilder,
};
pub trait ScalarFunctionAdapter: Sized + 'static {
unsafe extern "C" fn scalar_function_wrapper(
info: duckdb_function_info,
input: duckdb_data_chunk,
output: duckdb_vector,
) {
let info: ScalarFunctionInfo = unsafe { ScalarFunctionInfo::new(info) };
let extra = unsafe { erased_extra_info(&info) };
duck_scalar_unwind(&info,|| {
let chunk: DataChunk = unsafe { DataChunk::from_raw(input) };
let mut readers = Self::Args::create_column_readers(&chunk);
let fixed_count = readers.len();
let has_varargs = Self::varargs_element_type().is_some();
if has_varargs {
for column_index in fixed_count..chunk.column_count() {
readers.push(Self::varargs_create_reader(&chunk, column_index));
}
}
let row_count = chunk.size();
let result = if has_varargs {
let mut output_vec: Vec<Option<Self::Output>> = Vec::with_capacity(row_count);
for row in 0..row_count {
match Self::apply_varargs(&readers, row, fixed_count) {
Ok(r) => output_vec.push(r),
Err(e) => {
info.set_error(e.as_str());
return;
}
}
}
Ok(Some(output_vec))
} else {
let rows: Vec<Option<Self::Args>> = (0..row_count)
.map(|row| Self::Args::read_columns(&readers, row))
.collect();
Self::apply_batch(rows, extra)
};
match result {
Ok(Some(results)) => {
if results.len() != row_count {
info.set_error(
duck_error(format!(
"{}: batch implementation returned {} rows for {} input rows",
Self::NAME,
results.len(),
row_count
))
.as_str(),
);
return;
}
Self::Output::write_batch(output, &vec_option_to_ref(&results));
}
Ok(None) => {
let nulls: Vec<Option<&Self::Output>> = vec![None; row_count];
Self::Output::write_batch(output, &nulls);
}
Err(e) => info.set_error(e.as_str()),
}
});
}
fn null_handling() -> NullHandling {
NullHandling::DefaultNullHandling
}
fn volatile() -> bool {
false
}
fn varargs_element_type() -> Option<LogicalType> {
None
}
fn varargs_create_reader(_chunk: &DataChunk, _column_index: usize) -> DuckValueReader {
unreachable!("varargs_create_reader is only used when varargs_element_type() returns Some")
}
fn apply_varargs(
_readers: &[DuckValueReader],
_row: usize,
_fixed_count: usize,
) -> DuckOptionResult<Self::Output> {
unreachable!("apply_varargs is only used when varargs_element_type() returns Some")
}
fn apply_batch(
rows: Vec<Option<Self::Args>>,
extra: Option<&DuckExtraInfo>,
) -> DuckOptionResult<Vec<Option<Self::Output>>> {
let mut output_vec: Vec<Option<Self::Output>> = Vec::with_capacity(rows.len());
for args in rows {
output_vec.push(Self::apply_with_extra(args, extra)?);
}
Ok(Some(output_vec))
}
fn scalar_function_builder() -> ScalarFunctionBuilder {
let mut builder = ScalarFunctionBuilder::new(Self::NAME)
.function(Self::scalar_function_wrapper)
.null_handling(Self::null_handling())
.returns_logical(Self::Output::logical_type())
.with_params(Self::Args::column_types());
if Self::volatile() {
builder = set_volatile(builder);
}
if let Some(varargs_type) = Self::varargs_element_type() {
builder = set_varargs(builder, varargs_type);
}
if let Some((ptr, destroy)) = raw_extra_info(Self::extra_info()) {
builder = unsafe { builder.extra_info(ptr, destroy) };
}
builder
}
fn scalar_overload_builder() -> ScalarOverloadBuilder {
let mut builder = ScalarOverloadBuilder::new()
.function(Self::scalar_function_wrapper)
.null_handling(Self::null_handling())
.returns_logical(Self::Output::logical_type())
.with_params(Self::Args::column_types());
if let Some((ptr, destroy)) = raw_extra_info(Self::extra_info()) {
builder = unsafe { builder.extra_info(ptr, destroy) };
}
builder
}
unsafe fn register(con: duckdb_connection) -> DuckResult<()> {
unsafe { Self::scalar_function_builder().register(con) }
}
const NAME: &'static str;
type Args: DuckColumns;
type Output: DuckValueType;
fn extra_info() -> Option<DuckExtraInfo> {
None
}
fn apply_with_null(args_option: Option<Self::Args>) -> DuckOptionResult<Self::Output> {
if let Some(args) = args_option {
Self::apply(args)
} else {
Ok(None)
}
}
fn apply_with_extra(
args: Option<Self::Args>,
extra: Option<&DuckExtraInfo>,
) -> DuckOptionResult<Self::Output> {
let _ = extra;
Self::apply_with_null(args)
}
fn apply(args: Self::Args) -> DuckOptionResult<Self::Output>;
}
fn set_volatile(builder: ScalarFunctionBuilder) -> ScalarFunctionBuilder {
#[cfg(feature = "duckdb-1-5")]
{
builder.volatile()
}
#[cfg(not(feature = "duckdb-1-5"))]
{
builder
}
}
fn set_varargs(builder: ScalarFunctionBuilder, varargs_type: LogicalType) -> ScalarFunctionBuilder {
#[cfg(feature = "duckdb-1-5")]
{
builder.varargs_logical(varargs_type)
}
#[cfg(not(feature = "duckdb-1-5"))]
{
let _ = varargs_type;
builder
}
}