use libduckdb_sys::{
duckdb_data_chunk, duckdb_function_info, duckdb_scalar_function_get_extra_info, duckdb_scalar_function_set_error,
duckdb_vector,
};
use std::ffi::CStr;
use crate::{
Connection,
arrow_interop::WritableVector,
callback::{CallbackErrorSink, contain_callback},
core::{DataChunkHandle, LogicalTypeHandle},
inner_connection::InnerConnection,
};
use self::function::{ScalarFunction, ScalarFunctionSet};
mod function;
#[cfg(feature = "vscalar-arrow")]
pub mod arrow;
#[cfg(feature = "vscalar-arrow")]
pub use arrow::{ArrowFunctionSignature, ArrowScalarParams, VArrowScalar};
pub trait VScalar: Sized {
type State: Sized + Send + Sync + 'static;
fn invoke(
state: &Self::State,
input: &mut DataChunkHandle,
output: &mut dyn WritableVector,
) -> Result<(), Box<dyn std::error::Error>>;
fn signatures() -> Vec<ScalarFunctionSignature>;
fn volatile() -> bool {
false
}
}
pub enum ScalarParams {
Exact(Vec<LogicalTypeHandle>),
Variadic(LogicalTypeHandle),
}
pub struct ScalarFunctionSignature {
parameters: Option<ScalarParams>,
return_type: LogicalTypeHandle,
}
impl ScalarFunctionSignature {
pub fn exact(params: Vec<LogicalTypeHandle>, return_type: LogicalTypeHandle) -> Self {
Self {
parameters: Some(ScalarParams::Exact(params)),
return_type,
}
}
pub fn variadic(param: LogicalTypeHandle, return_type: LogicalTypeHandle) -> Self {
Self {
parameters: Some(ScalarParams::Variadic(param)),
return_type,
}
}
}
impl ScalarFunctionSignature {
pub(crate) fn register_with_scalar(&self, f: &ScalarFunction) {
f.set_return_type(&self.return_type);
match &self.parameters {
Some(ScalarParams::Exact(params)) => {
for param in params.iter() {
f.add_parameter(param);
}
}
Some(ScalarParams::Variadic(param)) => {
f.add_variadic_parameter(param);
}
None => {
}
}
}
}
#[derive(Debug)]
struct ScalarFunctionInfo(duckdb_function_info);
impl From<duckdb_function_info> for ScalarFunctionInfo {
fn from(ptr: duckdb_function_info) -> Self {
Self(ptr)
}
}
impl ScalarFunctionInfo {
pub unsafe fn get_extra_info<T>(&self) -> &T {
unsafe { &*(duckdb_scalar_function_get_extra_info(self.0).cast()) }
}
}
impl CallbackErrorSink for ScalarFunctionInfo {
fn set_c_error(&self, error: &CStr) {
unsafe { duckdb_scalar_function_set_error(self.0, error.as_ptr()) };
}
}
unsafe extern "C" fn scalar_func<T>(info: duckdb_function_info, input: duckdb_data_chunk, mut output: duckdb_vector)
where
T: VScalar,
{
let info = ScalarFunctionInfo::from(info);
contain_callback(&info, || unsafe {
let mut input = DataChunkHandle::new_unowned(input);
T::invoke(info.get_extra_info(), &mut input, &mut output)
});
}
impl Connection {
#[inline]
pub fn register_scalar_function<S: VScalar>(&self, name: &str) -> crate::Result<()>
where
S::State: Default,
{
let set = ScalarFunctionSet::new(name);
for signature in S::signatures() {
let scalar_function = ScalarFunction::new(name)?;
signature.register_with_scalar(&scalar_function);
scalar_function.set_function(Some(scalar_func::<S>));
if S::volatile() {
scalar_function.set_volatile();
}
scalar_function.set_extra_info(S::State::default());
set.add_function(scalar_function)?;
}
self.db.borrow_mut().register_scalar_function_set(set)
}
#[inline]
pub fn register_scalar_function_with_state<S: VScalar>(&self, name: &str, state: &S::State) -> crate::Result<()>
where
S::State: Clone,
{
let set = ScalarFunctionSet::new(name);
for signature in S::signatures() {
let scalar_function = ScalarFunction::new(name)?;
signature.register_with_scalar(&scalar_function);
scalar_function.set_function(Some(scalar_func::<S>));
if S::volatile() {
scalar_function.set_volatile();
}
scalar_function.set_extra_info(state.clone());
set.add_function(scalar_function)?;
}
self.db.borrow_mut().register_scalar_function_set(set)
}
}
impl InnerConnection {
pub fn register_scalar_function_set(&mut self, f: ScalarFunctionSet) -> crate::Result<()> {
f.register_with_connection(self.con)
}
}
#[cfg(test)]
mod tests;