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, erased_extra_info, raw_extra_info, vec_option_to_ref,
DuckOptionResult, DuckResult,
};
use libduckdb_sys::{
duckdb_aggregate_state, duckdb_connection, duckdb_data_chunk, duckdb_function_info,
duckdb_vector, idx_t,
};
use quack_rs::aggregate::{
AggregateFunctionBuilder, AggregateOverloadBuilder, AggregateState, FfiState,
};
use quack_rs::data_chunk::DataChunk;
use quack_rs::prelude::{AggregateFunctionInfo, NullHandling};
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: AggregateOverloadBuilder) -> AggregateOverloadBuilder {
let mut builder = 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())
.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::aggregate_function_builder().register(con) }
}
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")
}
}