mod function;
use std::ffi::CString;
use crate::{
connection::Connection,
error::Result,
ffi::{duckdb_data_chunk, duckdb_function_info, duckdb_vector},
};
use self::function::{ScalarFunction, ScalarFunctionInfo, ScalarFunctionSet};
use super::{
callback::contain_callback, data_chunk::DataChunkHandle, logical_type::LogicalType,
vector::VectorMut,
};
pub trait VScalar: Sized {
type State: Send + Sync + 'static;
fn signatures() -> Result<Vec<ScalarSignature>>;
fn invoke(
state: &Self::State,
input: &DataChunkHandle,
output: &mut VectorMut<'_>,
) -> super::UdfResult<()>;
fn volatile() -> bool {
false
}
fn special_handling() -> bool {
false
}
}
enum ScalarParams {
Exact(Vec<LogicalType>),
Variadic(LogicalType),
}
pub struct ScalarSignature {
parameters: Option<ScalarParams>,
return_type: LogicalType,
}
impl ScalarSignature {
pub fn exact(
parameters: Vec<LogicalType>,
return_type: LogicalType,
) -> Self {
Self { parameters: Some(ScalarParams::Exact(parameters)), return_type }
}
pub fn variadic(
parameter: LogicalType,
return_type: LogicalType,
) -> Self {
Self { parameters: Some(ScalarParams::Variadic(parameter)), return_type }
}
fn apply(
&self,
f: &ScalarFunction,
) {
f.set_return_type(&self.return_type);
match &self.parameters {
Some(ScalarParams::Exact(params)) => {
for p in params {
f.add_parameter(p);
}
},
Some(ScalarParams::Variadic(p)) => f.set_varargs(p),
None => {},
}
}
}
unsafe extern "C" fn scalar_trampoline<S: VScalar>(
info: duckdb_function_info,
input: duckdb_data_chunk,
output: duckdb_vector,
) {
let sink = ScalarFunctionInfo::from(info);
contain_callback(&sink, || {
let chunk = unsafe { DataChunkHandle::borrowed(input) };
let mut out = unsafe { VectorMut::new(output) };
let state = unsafe { sink.state::<S::State>() };
S::invoke(state, &chunk, &mut out)
});
}
impl Connection {
pub fn register_scalar_function<S: VScalar>(
&mut self,
name: &str,
) -> Result<()>
where
S::State: Default,
{
register_scalar_function_impl::<S>(self, name, S::State::default)
}
pub fn register_scalar_function_with_state<S: VScalar>(
&mut self,
name: &str,
state: S::State,
) -> Result<()>
where
S::State: Clone,
{
register_scalar_function_impl::<S>(self, name, move || state.clone())
}
}
fn register_scalar_function_impl<S: VScalar>(
conn: &mut Connection,
name: &str,
mut make_state: impl FnMut() -> S::State,
) -> Result<()> {
let c_name = CString::new(name)?;
let set = ScalarFunctionSet::new(&c_name);
for signature in S::signatures()? {
let f = ScalarFunction::new(&c_name);
signature.apply(&f);
f.set_function(Some(scalar_trampoline::<S>));
if S::volatile() {
f.set_volatile();
}
if S::special_handling() {
f.set_special_handling();
}
f.set_extra_info(make_state());
set.add_function(&f)?;
}
set.register(conn.raw_con(), name)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::connection::Connection;
struct AddOne;
impl VScalar for AddOne {
type State = ();
fn signatures() -> Result<Vec<ScalarSignature>> {
Ok(vec![ScalarSignature::exact(
vec![LogicalType::of::<i32>()?],
LogicalType::of::<i32>()?,
)])
}
fn invoke(
_state: &(),
input: &DataChunkHandle,
output: &mut VectorMut<'_>,
) -> super::super::UdfResult<()> {
let col = input.vector(0)?;
for row in 0..input.len() {
let v: i32 = col.get(row)?;
output.set(row, v + 1)?;
}
Ok(())
}
}
#[test]
fn register_and_call_scalar_function() {
let mut conn = Connection::open_in_memory().unwrap();
conn.register_scalar_function::<AddOne>("add_one").unwrap();
conn.execute_batch("CREATE TABLE t (v INTEGER)").unwrap();
conn.execute_batch("INSERT INTO t VALUES (1), (2), (41)").unwrap();
let result = conn.execute("SELECT add_one(v) AS r FROM t ORDER BY v").unwrap();
let rows: Vec<_> = result.collect::<Result<_>>().unwrap();
assert_eq!(rows.len(), 3);
assert_eq!(rows[2].get("r"), Some(&crate::types::value::DuckValue::Int(42)));
}
struct AlwaysFails;
impl VScalar for AlwaysFails {
type State = ();
fn signatures() -> Result<Vec<ScalarSignature>> {
Ok(vec![ScalarSignature::exact(
vec![LogicalType::of::<i32>()?],
LogicalType::of::<i32>()?,
)])
}
fn invoke(
_state: &(),
_input: &DataChunkHandle,
_output: &mut VectorMut<'_>,
) -> super::super::UdfResult<()> {
Err("deliberate failure".into())
}
}
#[test]
fn invoke_error_surfaces_as_query_error_and_connection_stays_usable() {
let mut conn = Connection::open_in_memory().unwrap();
conn.register_scalar_function::<AlwaysFails>("always_fails").unwrap();
conn.execute_batch("CREATE TABLE t (v INTEGER)").unwrap();
conn.execute_batch("INSERT INTO t VALUES (1)").unwrap();
let err = match conn.execute("SELECT always_fails(v) FROM t") {
Ok(_) => panic!("expected an error"),
Err(e) => e,
};
assert!(err.to_string().contains("deliberate failure"), "{err}");
conn.execute_batch("INSERT INTO t VALUES (2)").unwrap();
}
struct AlwaysPanics;
impl VScalar for AlwaysPanics {
type State = ();
fn signatures() -> Result<Vec<ScalarSignature>> {
Ok(vec![ScalarSignature::exact(
vec![LogicalType::of::<i32>()?],
LogicalType::of::<i32>()?,
)])
}
fn invoke(
_state: &(),
_input: &DataChunkHandle,
_output: &mut VectorMut<'_>,
) -> super::super::UdfResult<()> {
panic!("deliberate panic")
}
}
#[test]
#[cfg(panic = "unwind")]
fn invoke_panic_is_contained_and_connection_stays_usable() {
let mut conn = Connection::open_in_memory().unwrap();
conn.register_scalar_function::<AlwaysPanics>("always_panics").unwrap();
conn.execute_batch("CREATE TABLE t (v INTEGER)").unwrap();
conn.execute_batch("INSERT INTO t VALUES (1)").unwrap();
let err = match conn.execute("SELECT always_panics(v) FROM t") {
Ok(_) => panic!("expected an error"),
Err(e) => e,
};
assert!(err.to_string().contains("deliberate panic"), "{err}");
conn.execute_batch("INSERT INTO t VALUES (2)").unwrap();
}
}