use crate::duck_columns::DuckColumns;
use crate::utils::builder_with_params::BuilderWithParams;
use crate::value_types::duck_value_type::DuckValueType;
use crate::{
DuckExtraInfo, duck_aggregate_unwind, duck_error, erased_extra_info, raw_extra_info,
vec_option_to_ref, DuckOptionResult, DuckResult,
};
use libduckdb_sys::{
DuckDBSuccess, duckdb_add_aggregate_function_to_set, duckdb_aggregate_function,
duckdb_aggregate_function_add_parameter, duckdb_aggregate_function_set,
duckdb_aggregate_function_set_destructor, duckdb_aggregate_function_set_extra_info,
duckdb_aggregate_function_set_functions, duckdb_aggregate_function_set_name,
duckdb_aggregate_function_set_return_type, duckdb_aggregate_function_set_special_handling,
duckdb_aggregate_state, duckdb_connection, duckdb_create_aggregate_function,
duckdb_create_aggregate_function_set, duckdb_data_chunk, duckdb_destroy_aggregate_function,
duckdb_destroy_aggregate_function_set, duckdb_function_info,
duckdb_register_aggregate_function_set, duckdb_vector, idx_t,
};
use quack_rs::aggregate::builder::OverloadBuilder;
use quack_rs::aggregate::{AggregateFunctionBuilder, AggregateState, FfiState};
use quack_rs::data_chunk::DataChunk;
use quack_rs::error::ExtensionError;
use quack_rs::prelude::{AggregateFunctionInfo, NullHandling};
use std::ffi::CString;
pub trait AggregateFunctionAdapter: AggregateState + Sized + 'static {
unsafe extern "C" fn c_state_size(_info: duckdb_function_info) -> idx_t {
unsafe { FfiState::<Self>::size_callback(_info) }
}
unsafe extern "C" fn c_state_init(info: duckdb_function_info, state: duckdb_aggregate_state) {
unsafe { FfiState::<Self>::init_callback(info, state) };
}
unsafe extern "C" fn c_update(
info: duckdb_function_info,
input: duckdb_data_chunk,
states: *mut duckdb_aggregate_state,
) {
let info = unsafe { AggregateFunctionInfo::new(info) };
let extra = unsafe { erased_extra_info(&info) };
duck_aggregate_unwind(&info, || {
let chunk = unsafe { DataChunk::from_raw(input) };
let readers = Self::Args::create_column_readers(&chunk);
let row_count = chunk.size();
for row in 0..row_count {
let args = Self::Args::read_columns(&readers, row);
let state_ptr = unsafe { *states.add(row) };
if let Some(st) = unsafe { FfiState::<Self>::with_state_mut(state_ptr) } {
let result = st.handle_row_with_extra(args, extra);
if let Err(e) = result {
info.set_error(e.as_str());
return;
}
}
}
});
}
unsafe extern "C" fn c_combine(
info: duckdb_function_info,
source: *mut duckdb_aggregate_state,
target: *mut duckdb_aggregate_state,
count: idx_t,
) {
let info = unsafe { AggregateFunctionInfo::new(info) };
let extra = unsafe { erased_extra_info(&info) };
duck_aggregate_unwind(&info, || {
for i in 0..count as usize {
let src_ptr = unsafe { *source.add(i) };
let tgt_ptr = unsafe { *target.add(i) };
let src = unsafe { FfiState::<Self>::with_state(src_ptr) };
let tgt = unsafe { FfiState::<Self>::with_state_mut(tgt_ptr) };
if let (Some(s), Some(t)) = (src, tgt) {
let result = t.combine_with_extra(s, extra);
if let Err(e) = result {
info.set_error(e.as_str());
return;
}
}
}
});
}
unsafe extern "C" fn c_finalize(
info: duckdb_function_info,
source: *mut duckdb_aggregate_state,
result: duckdb_vector,
count: idx_t,
offset: idx_t,
) {
let info = unsafe { AggregateFunctionInfo::new(info) };
let extra = unsafe { erased_extra_info(&info) };
duck_aggregate_unwind(&info, || {
if offset != 0 {
info.set_error(
format!(
"non-zero aggregate finalize result offset is not supported: {}",
offset
)
.as_str(),
);
return;
}
let mut output_vec: Vec<Option<Self::Output>> = Vec::with_capacity(count as usize);
for i in 0..count as usize {
let state_ptr = unsafe { *source.add(i) };
match unsafe { FfiState::<Self>::with_state(state_ptr) } {
Some(st) => {
let result1 = st.result_with_extra(extra);
match result1 {
Ok(r) => output_vec.push(r),
Err(e) => {
info.set_error(e.as_str());
return;
}
};
}
None => output_vec.push(None),
}
}
Self::Output::write_batch(result, &vec_option_to_ref(&output_vec));
});
}
unsafe extern "C" fn c_state_destroy(states: *mut duckdb_aggregate_state, count: idx_t) {
unsafe { FfiState::<Self>::destroy_callback(states, count) };
}
fn null_handling() -> NullHandling {
NullHandling::DefaultNullHandling
}
fn extra_info() -> Option<DuckExtraInfo> {
None
}
fn aggregate_function_builder() -> AggregateFunctionBuilder {
let mut builder = AggregateFunctionBuilder::new(Self::NAME)
.state_size(Self::c_state_size)
.init(Self::c_state_init)
.update(Self::c_update)
.combine(Self::c_combine)
.finalize(Self::c_finalize)
.destructor(Self::c_state_destroy)
.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
}
fn aggregate_overload_builder(builder: OverloadBuilder) -> OverloadBuilder {
builder
.state_size(Self::c_state_size)
.init(Self::c_state_init)
.update(Self::c_update)
.combine(Self::c_combine)
.finalize(Self::c_finalize)
.destructor(Self::c_state_destroy)
.null_handling(Self::null_handling())
.with_params(Self::Args::column_types())
}
unsafe fn register(con: duckdb_connection) -> DuckResult<()> {
unsafe { Self::aggregate_function_builder().register(con) }
}
fn create_aggregate_function_guard(name: &CString) -> AggregateFunctionGuard {
let func = unsafe { duckdb_create_aggregate_function() };
unsafe { duckdb_aggregate_function_set_name(func, name.as_ptr()) };
for lt in Self::Args::column_types() {
unsafe { duckdb_aggregate_function_add_parameter(func, lt.as_raw()) };
}
unsafe {
duckdb_aggregate_function_set_return_type(func, Self::Output::logical_type().as_raw())
};
unsafe {
duckdb_aggregate_function_set_functions(
func,
Some(Self::c_state_size),
Some(Self::c_state_init),
Some(Self::c_update),
Some(Self::c_combine),
Some(Self::c_finalize),
)
};
unsafe { duckdb_aggregate_function_set_destructor(func, Some(Self::c_state_destroy)) };
if Self::null_handling() == NullHandling::SpecialNullHandling {
unsafe { duckdb_aggregate_function_set_special_handling(func) };
}
if let Some((ptr, destroy)) = raw_extra_info(Self::extra_info()) {
unsafe { duckdb_aggregate_function_set_extra_info(func, ptr, destroy) };
}
AggregateFunctionGuard {
name: Self::NAME.to_string(),
c_agg: func,
}
}
fn aggregate_function_guard() -> AggregateFunctionGuard {
let name = CString::new(Self::NAME).expect("function name must not contain null bytes");
Self::create_aggregate_function_guard(&name)
}
const NAME: &'static str;
type Args: DuckColumns;
type Output: DuckValueType;
fn handle_row_with_null(&mut self, args: Option<Self::Args>) -> DuckResult<()> {
if let Some(args) = args {
self.handle_row(args)?;
}
Ok(())
}
fn handle_row_with_extra(
&mut self,
args: Option<Self::Args>,
extra: Option<&DuckExtraInfo>,
) -> DuckResult<()> {
let _ = extra;
self.handle_row_with_null(args)
}
fn handle_row(&mut self, args: Self::Args) -> DuckResult<()>;
fn combine(&mut self, other: &Self) -> DuckResult<()>;
fn result(&self) -> DuckOptionResult<Self::Output>;
fn combine_with_extra(
&mut self,
other: &Self,
extra: Option<&DuckExtraInfo>,
) -> DuckResult<()> {
let _ = extra;
self.combine(other)
}
fn result_with_extra(&self, extra: Option<&DuckExtraInfo>) -> DuckOptionResult<Self::Output> {
let _ = extra;
self.result()
}
}
pub trait DuckAggregateState {
type Output: DuckValueType;
fn combine(&mut self, other: &Self) -> DuckResult<()> {
self.simple_combine(other);
Ok(())
}
fn simple_combine(&mut self, _other: &Self) {
todo!("simple_combine is not implemented")
}
fn result(&self) -> DuckOptionResult<Self::Output> {
Ok(Some(self.simple_result()))
}
fn simple_result(&self) -> Self::Output {
todo!("simple_result is not implemented")
}
}
pub struct AggregateFunctionGuard {
name: String,
c_agg: duckdb_aggregate_function,
}
impl AggregateFunctionGuard {
pub fn as_raw(&self) -> duckdb_aggregate_function {
self.c_agg
}
}
impl Drop for AggregateFunctionGuard {
fn drop(&mut self) {
unsafe { duckdb_destroy_aggregate_function(&raw mut self.c_agg) };
}
}
pub struct AggregateFunctionSetGuard {
c_set: duckdb_aggregate_function_set,
}
impl AggregateFunctionSetGuard {
fn add_overload(&mut self, func: AggregateFunctionGuard) -> DuckResult<()> {
let result = unsafe { duckdb_add_aggregate_function_to_set(self.as_raw(), func.as_raw()) };
if result != DuckDBSuccess {
return Err(duck_error(format!(
"duckdb_add_aggregate_function_to_set failed: {}",
func.name
)));
}
Ok(())
}
pub fn as_raw(&self) -> duckdb_aggregate_function_set {
self.c_set
}
}
impl Drop for AggregateFunctionSetGuard {
fn drop(&mut self) {
unsafe { duckdb_destroy_aggregate_function_set(&raw mut self.c_set) };
}
}
pub struct DuckfnAggregateFunctionSetBuilder {
pub name: CString,
pub overloads: Vec<AggregateFunctionGuard>,
}
impl DuckfnAggregateFunctionSetBuilder {
pub fn new(name: &str, overloads: Vec<AggregateFunctionGuard>) -> Self {
Self {
name: CString::new(name).expect("function name must not contain null bytes"),
overloads,
}
}
pub unsafe fn register(self, con: duckdb_connection) -> DuckResult<()> {
let mut set = AggregateFunctionSetGuard {
c_set: unsafe { duckdb_create_aggregate_function_set(self.name.as_ptr()) },
};
for x in self.overloads {
set.add_overload(x)?;
}
let result = unsafe { duckdb_register_aggregate_function_set(con, set.as_raw()) };
if result != DuckDBSuccess {
return Err(ExtensionError::new(format!(
"duckdb_register_aggregate_function_set failed for '{}'",
self.name.to_string_lossy()
)));
}
Ok(())
}
}