#![forbid(unsafe_code)]
use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use arrow::array::{ArrayRef, Int64Array};
use arrow::datatypes::{DataType, Field, Int64Type, Schema};
use arrow::record_batch::RecordBatch;
#[derive(Debug, thiserror::Error)]
pub enum UdfError {
#[error("Arrow error: {0}")]
Arrow(String),
#[error("Execution error: {message}")]
Execution { message: String },
#[error("Panic: {0}")]
Panic(String),
#[error("Invalid argument: {message}")]
InvalidArgument { message: String },
}
impl From<arrow::error::ArrowError> for UdfError {
fn from(e: arrow::error::ArrowError) -> Self {
UdfError::Arrow(e.to_string())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum Volatility {
#[default]
Immutable,
Stable,
Volatile,
}
pub trait ScalarUdf: Send + Sync + fmt::Debug {
fn name(&self) -> &str;
fn input_schema(&self) -> &Schema;
fn output_field(&self) -> &Field;
fn volatility(&self) -> Volatility {
Volatility::Immutable
}
fn call(&self, batch: &RecordBatch) -> Result<ArrayRef, UdfError>;
}
#[derive(Debug, Default, Clone)]
pub struct AggState {
pub data: Vec<u8>,
}
#[derive(Debug, Clone)]
pub enum ScalarValue {
Null,
Int64(i64),
Float64(f64),
Utf8(String),
Boolean(bool),
Bytes(Vec<u8>),
}
pub trait AggregateUdf: Send + Sync + fmt::Debug {
fn name(&self) -> &str;
fn input_schema(&self) -> &Schema;
fn output_field(&self) -> &Field;
fn volatility(&self) -> Volatility {
Volatility::Immutable
}
fn accumulate(&self, state: &mut AggState, batch: &RecordBatch) -> Result<(), UdfError>;
fn finalize(&self, state: AggState) -> Result<ScalarValue, UdfError>;
fn merge(&self, a: AggState, b: AggState) -> Result<AggState, UdfError>;
}
pub trait TableUdf: Send + Sync + fmt::Debug {
fn name(&self) -> &str;
fn output_schema(&self) -> &Schema;
fn call(&self, args: &[ScalarValue]) -> Result<RecordBatch, UdfError>;
}
pub trait CoGroupUdf: Send + Sync + fmt::Debug {
fn name(&self) -> &str;
fn left_schema(&self) -> &Schema;
fn right_schema(&self) -> &Schema;
fn output_schema(&self) -> &Schema;
fn call(
&self,
key: &str,
left: &[RecordBatch],
right: &[RecordBatch],
) -> Result<Vec<RecordBatch>, UdfError>;
}
pub trait MapPandasIterUdf: Send + Sync + fmt::Debug {
fn name(&self) -> &str;
fn input_schema(&self) -> &Schema;
fn output_schema(&self) -> &Schema;
fn map_batches(&self, batches: &[RecordBatch]) -> Result<Vec<RecordBatch>, UdfError>;
}
#[derive(Debug, Default)]
pub struct UdfRegistry {
scalars: HashMap<String, Arc<dyn ScalarUdf>>,
aggregates: HashMap<String, Arc<dyn AggregateUdf>>,
tables: HashMap<String, Arc<dyn TableUdf>>,
co_groups: HashMap<String, Arc<dyn CoGroupUdf>>,
map_pandas_iters: HashMap<String, Arc<dyn MapPandasIterUdf>>,
}
impl UdfRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register_scalar(&mut self, udf: Arc<dyn ScalarUdf>) {
self.scalars.insert(udf.name().to_owned(), udf);
}
pub fn remove_scalar(&mut self, name: &str) -> Option<Arc<dyn ScalarUdf>> {
self.scalars.remove(name)
}
pub fn register_aggregate(&mut self, udf: Arc<dyn AggregateUdf>) {
self.aggregates.insert(udf.name().to_owned(), udf);
}
pub fn register_table(&mut self, udf: Arc<dyn TableUdf>) {
self.tables.insert(udf.name().to_owned(), udf);
}
pub fn register_co_group(&mut self, udf: Arc<dyn CoGroupUdf>) {
self.co_groups.insert(udf.name().to_owned(), udf);
}
pub fn register_map_pandas_iter(&mut self, udf: Arc<dyn MapPandasIterUdf>) {
self.map_pandas_iters.insert(udf.name().to_owned(), udf);
}
pub fn get_scalar(&self, name: &str) -> Option<&Arc<dyn ScalarUdf>> {
self.scalars.get(name)
}
pub fn get_aggregate(&self, name: &str) -> Option<&Arc<dyn AggregateUdf>> {
self.aggregates.get(name)
}
pub fn get_table(&self, name: &str) -> Option<&Arc<dyn TableUdf>> {
self.tables.get(name)
}
pub fn get_co_group(&self, name: &str) -> Option<&Arc<dyn CoGroupUdf>> {
self.co_groups.get(name)
}
pub fn get_map_pandas_iter(&self, name: &str) -> Option<&Arc<dyn MapPandasIterUdf>> {
self.map_pandas_iters.get(name)
}
pub fn scalar_names(&self) -> Vec<&str> {
let mut names: Vec<&str> = self.scalars.keys().map(String::as_str).collect();
names.sort_unstable();
names
}
pub fn aggregate_names(&self) -> Vec<&str> {
let mut names: Vec<&str> = self.aggregates.keys().map(String::as_str).collect();
names.sort_unstable();
names
}
pub fn table_names(&self) -> Vec<&str> {
let mut names: Vec<&str> = self.tables.keys().map(String::as_str).collect();
names.sort_unstable();
names
}
pub fn co_group_names(&self) -> Vec<&str> {
let mut names: Vec<&str> = self.co_groups.keys().map(String::as_str).collect();
names.sort_unstable();
names
}
pub fn map_pandas_iter_names(&self) -> Vec<&str> {
let mut names: Vec<&str> = self.map_pandas_iters.keys().map(String::as_str).collect();
names.sort_unstable();
names
}
pub fn execute_scalar_with_limits(
&self,
name: &str,
batch: &RecordBatch,
limits: &ResourceLimits,
executor: &dyn SandboxedUdfExecutor,
) -> Result<ArrayRef, UdfError> {
let udf = self
.get_scalar(name)
.ok_or_else(|| UdfError::InvalidArgument {
message: format!("unknown scalar UDF: {}", name),
})?;
executor.execute_with_limits(udf.as_ref(), batch, limits)
}
}
#[derive(Debug)]
pub struct MultiplyScalarUdf {
name: String,
column: String,
factor: i64,
input_schema: Schema,
output_field: Field,
}
impl MultiplyScalarUdf {
pub fn new(name: impl Into<String>, column: impl Into<String>, factor: i64) -> Self {
let column: String = column.into();
let input_schema = Schema::new(vec![Field::new(column.clone(), DataType::Int64, true)]);
let output_field = Field::new("result", DataType::Int64, true);
Self {
name: name.into(),
column,
factor,
input_schema,
output_field,
}
}
}
impl ScalarUdf for MultiplyScalarUdf {
fn name(&self) -> &str {
&self.name
}
fn input_schema(&self) -> &Schema {
&self.input_schema
}
fn output_field(&self) -> &Field {
&self.output_field
}
fn call(&self, batch: &RecordBatch) -> Result<ArrayRef, UdfError> {
let col_idx =
batch
.schema()
.index_of(&self.column)
.map_err(|_| UdfError::InvalidArgument {
message: format!("column '{}' not found in batch", self.column),
})?;
let array = batch.column(col_idx);
let int_array = array.as_any().downcast_ref::<Int64Array>().ok_or_else(|| {
UdfError::InvalidArgument {
message: format!("column '{}' is not Int64", self.column),
}
})?;
let factor = self.factor;
let result =
arrow::compute::kernels::arity::unary::<Int64Type, _, Int64Type>(int_array, |x| {
x.wrapping_mul(factor)
});
Ok(Arc::new(result))
}
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::{Array, Int64Array};
use arrow::datatypes::{DataType, Field, Schema};
use arrow::record_batch::RecordBatch;
use std::sync::Arc;
fn read_i64_state(state: &AggState) -> i64 {
if state.data.len() == 8 {
let mut buf = [0u8; 8];
buf.copy_from_slice(&state.data[..8]);
i64::from_le_bytes(buf)
} else {
0
}
}
#[derive(Debug)]
struct SumAggUdf {
input_schema: Schema,
output_field: Field,
}
impl SumAggUdf {
fn new() -> Self {
let input_schema = Schema::new(vec![Field::new("value", DataType::Int64, true)]);
let output_field = Field::new("sum", DataType::Int64, false);
Self {
input_schema,
output_field,
}
}
}
impl AggregateUdf for SumAggUdf {
fn name(&self) -> &str {
"sum_agg"
}
fn input_schema(&self) -> &Schema {
&self.input_schema
}
fn output_field(&self) -> &Field {
&self.output_field
}
fn accumulate(&self, state: &mut AggState, batch: &RecordBatch) -> Result<(), UdfError> {
let col = batch
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.ok_or_else(|| UdfError::InvalidArgument {
message: "expected Int64".into(),
})?;
let mut current: i64 = read_i64_state(state);
for v in col.iter().flatten() {
current += v;
}
state.data = current.to_le_bytes().to_vec();
Ok(())
}
fn finalize(&self, state: AggState) -> Result<ScalarValue, UdfError> {
Ok(ScalarValue::Int64(read_i64_state(&state)))
}
fn merge(&self, a: AggState, b: AggState) -> Result<AggState, UdfError> {
Ok(AggState {
data: (read_i64_state(&a) + read_i64_state(&b))
.to_le_bytes()
.to_vec(),
})
}
}
#[derive(Debug)]
struct ConstantTableUdf {
schema: Schema,
value: i64,
}
impl ConstantTableUdf {
fn new(value: i64) -> Self {
let schema = Schema::new(vec![Field::new("constant", DataType::Int64, false)]);
Self { schema, value }
}
}
impl TableUdf for ConstantTableUdf {
fn name(&self) -> &str {
"constant_table"
}
fn output_schema(&self) -> &Schema {
&self.schema
}
fn call(&self, _args: &[ScalarValue]) -> Result<RecordBatch, UdfError> {
let array = Int64Array::from(vec![self.value]);
RecordBatch::try_new(Arc::new(self.schema.clone()), vec![Arc::new(array)])
.map_err(UdfError::from)
}
}
#[test]
fn scalar_udf_registry_round_trip() {
let mut registry = UdfRegistry::new();
let udf = Arc::new(MultiplyScalarUdf::new("double", "x", 2));
registry.register_scalar(udf);
let found = registry
.get_scalar("double")
.expect("UDF must be registered");
assert_eq!(found.name(), "double");
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(vec![1_i64, 2, 3]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).expect("valid batch");
let result = found.call(&batch).expect("call must succeed");
let result_array = result
.as_any()
.downcast_ref::<Int64Array>()
.expect("result must be Int64");
assert_eq!(result_array.len(), 3);
assert_eq!(result_array.value(0), 2);
assert_eq!(result_array.value(1), 4);
assert_eq!(result_array.value(2), 6);
}
#[test]
fn aggregate_udf_state_lifecycle() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let array = Int64Array::from(vec![10_i64, 20]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).expect("valid batch");
let mut state = AggState::default();
udf.accumulate(&mut state, &batch).expect("accumulate ok");
let result = udf.finalize(state).expect("finalize ok");
match result {
ScalarValue::Int64(v) => assert_eq!(v, 30),
other => panic!("unexpected ScalarValue: {other:?}"),
}
}
#[test]
fn udf_error_display() {
let e1 = UdfError::Arrow("bad array".to_owned());
assert!(e1.to_string().contains("Arrow error"));
assert!(e1.to_string().contains("bad array"));
let e2 = UdfError::Execution {
message: "runtime fault".to_owned(),
};
assert!(e2.to_string().contains("Execution error"));
assert!(e2.to_string().contains("runtime fault"));
let e3 = UdfError::Panic("thread panicked".to_owned());
assert!(e3.to_string().contains("Panic"));
assert!(e3.to_string().contains("thread panicked"));
let e4 = UdfError::InvalidArgument {
message: "wrong type".to_owned(),
};
assert!(e4.to_string().contains("Invalid argument"));
assert!(e4.to_string().contains("wrong type"));
}
#[test]
fn registry_scalar_names_returns_registered_names() {
let mut registry = UdfRegistry::new();
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("triple", "v", 3)));
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("quadruple", "v", 4)));
let names = registry.scalar_names();
assert_eq!(names.len(), 2);
assert!(names.contains(&"triple"));
assert!(names.contains(&"quadruple"));
}
#[test]
fn table_udf_produces_record_batch() {
let mut registry = UdfRegistry::new();
let udtf = Arc::new(ConstantTableUdf::new(42));
registry.register_table(udtf);
let found = registry
.get_table("constant_table")
.expect("UDTF must be registered");
let batch = found.call(&[]).expect("call must succeed");
assert_eq!(batch.num_rows(), 1);
assert_eq!(batch.schema().field(0).name(), "constant");
let col = batch
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.expect("Int64");
assert_eq!(col.value(0), 42);
}
#[test]
fn udaf_distributed_merge_matches_single_partition() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let partition_a = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![1_i64, 2, 3, 4]))],
)
.expect("valid partition_a batch");
let partition_b = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![5_i64, 6, 7]))],
)
.expect("valid partition_b batch");
let mut state_a = AggState::default();
udf.accumulate(&mut state_a, &partition_a)
.expect("accumulate partition_a");
let mut state_b = AggState::default();
udf.accumulate(&mut state_b, &partition_b)
.expect("accumulate partition_b");
let partial_a = udf
.finalize(AggState {
data: state_a.data.clone(),
})
.expect("finalize partial_a");
let partial_b = udf
.finalize(AggState {
data: state_b.data.clone(),
})
.expect("finalize partial_b");
assert!(
matches!(partial_a, ScalarValue::Int64(10)),
"partial sum of partition_a must be 10, got {partial_a:?}",
);
assert!(
matches!(partial_b, ScalarValue::Int64(18)),
"partial sum of partition_b must be 18, got {partial_b:?}",
);
let merged_state = udf.merge(state_a, state_b).expect("merge partial states");
let distributed_result = udf.finalize(merged_state).expect("finalize merged state");
let all_values = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![1_i64, 2, 3, 4, 5, 6, 7]))],
)
.expect("valid all-values batch");
let mut single_state = AggState::default();
udf.accumulate(&mut single_state, &all_values)
.expect("accumulate single partition");
let single_result = udf
.finalize(single_state)
.expect("finalize single-partition state");
assert!(
matches!(distributed_result, ScalarValue::Int64(28)),
"distributed merge must produce 28, got {distributed_result:?}",
);
assert!(
matches!(single_result, ScalarValue::Int64(28)),
"single-partition path must produce 28, got {single_result:?}",
);
let distributed_val = match distributed_result {
ScalarValue::Int64(v) => v,
other => panic!("expected Int64, got {other:?}"),
};
let single_val = match single_result {
ScalarValue::Int64(v) => v,
other => panic!("expected Int64, got {other:?}"),
};
assert_eq!(
distributed_val, single_val,
"distributed merge ({distributed_val}) must equal single-partition result ({single_val})",
);
}
#[test]
fn udaf_merge_with_empty_state_is_noop() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let partition = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![10_i64, 20, 30]))],
)
.expect("valid partition batch");
let mut non_empty_state = AggState::default();
udf.accumulate(&mut non_empty_state, &partition)
.expect("accumulate");
let merged_right = udf
.merge(
AggState {
data: non_empty_state.data.clone(),
},
AggState::default(),
)
.expect("merge with empty right");
let merged_left = udf
.merge(
AggState::default(),
AggState {
data: non_empty_state.data.clone(),
},
)
.expect("merge with empty left");
let result_right = udf.finalize(merged_right).expect("finalize right merge");
let result_left = udf.finalize(merged_left).expect("finalize left merge");
assert!(
matches!(result_right, ScalarValue::Int64(60)),
"merge with empty right must yield 60, got {result_right:?}",
);
assert!(
matches!(result_left, ScalarValue::Int64(60)),
"merge with empty left must yield 60, got {result_left:?}",
);
}
#[test]
fn udaf_merge_three_partitions() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let make_batch = |vals: Vec<i64>| {
RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(Int64Array::from(vals))])
.expect("valid batch")
};
let mut s1 = AggState::default();
let mut s2 = AggState::default();
let mut s3 = AggState::default();
udf.accumulate(&mut s1, &make_batch(vec![100]))
.expect("acc p1");
udf.accumulate(&mut s2, &make_batch(vec![200, 300]))
.expect("acc p2");
udf.accumulate(&mut s3, &make_batch(vec![400, 500, 600]))
.expect("acc p3");
let m12 = udf.merge(s1, s2).expect("merge s1+s2");
let m123 = udf.merge(m12, s3).expect("merge (s1+s2)+s3");
let result = udf.finalize(m123).expect("finalize three-partition merge");
assert!(
matches!(result, ScalarValue::Int64(2100)),
"three-partition merge must yield 2100, got {result:?}",
);
}
#[test]
fn multiply_scalar_negative_factor() {
let udf = MultiplyScalarUdf::new("neg", "x", -3);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(vec![2_i64, -5, 0]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.value(0), -6);
assert_eq!(arr.value(1), 15);
assert_eq!(arr.value(2), 0);
}
#[test]
fn multiply_scalar_zero_factor() {
let udf = MultiplyScalarUdf::new("zero", "x", 0);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(vec![100_i64, 200]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.value(0), 0);
assert_eq!(arr.value(1), 0);
}
#[test]
fn multiply_scalar_one_factor() {
let udf = MultiplyScalarUdf::new("id", "x", 1);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(vec![42_i64]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.value(0), 42);
}
#[test]
fn multiply_scalar_large_values() {
let udf = MultiplyScalarUdf::new("large", "x", 2);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(vec![i64::MAX / 2, i64::MIN / 2]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.value(0), i64::MAX / 2 * 2);
assert_eq!(arr.value(1), i64::MIN / 2 * 2);
}
#[test]
fn multiply_scalar_empty_batch() {
let udf = MultiplyScalarUdf::new("empty", "x", 5);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(Vec::<i64>::new());
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.len(), 0);
}
#[test]
fn multiply_scalar_column_not_found() {
let udf = MultiplyScalarUdf::new("m", "missing_col", 1);
let schema = Arc::new(Schema::new(vec![Field::new(
"other",
DataType::Int64,
true,
)]));
let array = Int64Array::from(vec![1_i64]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let err = udf.call(&batch).unwrap_err();
assert!(matches!(err, UdfError::InvalidArgument { .. }));
assert!(err.to_string().contains("missing_col"));
}
#[test]
fn multiply_scalar_wrong_type_column() {
let udf = MultiplyScalarUdf::new("m", "x", 1);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, true)]));
let array = arrow::array::StringArray::from(vec!["hello"]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let err = udf.call(&batch).unwrap_err();
assert!(matches!(err, UdfError::InvalidArgument { .. }));
assert!(err.to_string().contains("not Int64"));
}
#[test]
fn multiply_scalar_null_values() {
let udf = MultiplyScalarUdf::new("m", "x", 10);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let mut builder = arrow::array::Int64Builder::new();
builder.append_value(5);
builder.append_null();
builder.append_value(3);
let array = builder.finish();
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.value(0), 50);
assert!(arr.is_null(1));
assert_eq!(arr.value(2), 30);
}
#[test]
fn multiply_scalar_output_schema() {
let udf = MultiplyScalarUdf::new("m", "input", 2);
assert_eq!(udf.output_field().name(), "result");
assert_eq!(udf.output_field().data_type(), &DataType::Int64);
}
#[test]
fn multiply_scalar_input_schema() {
let udf = MultiplyScalarUdf::new("m", "my_col", 1);
let schema = udf.input_schema();
assert_eq!(schema.fields().len(), 1);
assert_eq!(schema.field(0).name(), "my_col");
}
#[test]
fn udf_registry_scalar_override() {
let mut registry = UdfRegistry::new();
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("f", "x", 2)));
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("f", "x", 3)));
let udf = registry.get_scalar("f").unwrap();
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let array = Int64Array::from(vec![1_i64]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.value(0), 3); }
#[test]
fn udf_registry_aggregate_override() {
let mut registry = UdfRegistry::new();
registry.register_aggregate(Arc::new(SumAggUdf::new()));
registry.register_aggregate(Arc::new(SumAggUdf::new()));
assert_eq!(registry.aggregate_names().len(), 1);
}
#[test]
fn udf_registry_table_override() {
let mut registry = UdfRegistry::new();
registry.register_table(Arc::new(ConstantTableUdf::new(1)));
registry.register_table(Arc::new(ConstantTableUdf::new(2)));
assert_eq!(registry.table_names().len(), 1);
let udf = registry.get_table("constant_table").unwrap();
let batch = udf.call(&[]).unwrap();
let col = batch
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
assert_eq!(col.value(0), 2);
}
#[test]
fn udf_registry_missing_scalar_returns_none() {
let registry = UdfRegistry::new();
assert!(registry.get_scalar("nonexistent").is_none());
}
#[test]
fn udf_registry_missing_aggregate_returns_none() {
let registry = UdfRegistry::new();
assert!(registry.get_aggregate("nonexistent").is_none());
}
#[test]
fn udf_registry_missing_table_returns_none() {
let registry = UdfRegistry::new();
assert!(registry.get_table("nonexistent").is_none());
}
#[test]
fn udf_registry_empty_names() {
let registry = UdfRegistry::new();
assert!(registry.scalar_names().is_empty());
assert!(registry.aggregate_names().is_empty());
assert!(registry.table_names().is_empty());
}
#[test]
fn udf_registry_multiple_scalars_sorted() {
let mut registry = UdfRegistry::new();
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("z", "x", 1)));
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("a", "x", 1)));
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("m", "x", 1)));
let names = registry.scalar_names();
assert_eq!(names, vec!["a", "m", "z"]);
}
#[test]
fn udf_registry_remove_scalar_returns_registration() {
let mut registry = UdfRegistry::new();
registry.register_scalar(Arc::new(MultiplyScalarUdf::new("double", "x", 2)));
let removed = registry
.remove_scalar("double")
.expect("registered scalar should be returned");
assert_eq!(removed.name(), "double");
assert!(registry.get_scalar("double").is_none());
}
#[test]
fn udf_registry_multiple_aggregates_sorted() {
let mut registry = UdfRegistry::new();
registry.register_aggregate(Arc::new(SumAggUdf::new()));
let names = registry.aggregate_names();
assert_eq!(names, vec!["sum_agg"]);
}
#[test]
fn udf_registry_multiple_tables_sorted() {
let mut registry = UdfRegistry::new();
registry.register_table(Arc::new(ConstantTableUdf::new(1)));
let names = registry.table_names();
assert_eq!(names, vec!["constant_table"]);
}
#[test]
fn aggregate_empty_batch_finalize() {
let udf = SumAggUdf::new();
let state = AggState::default();
let result = udf.finalize(state).unwrap();
assert!(matches!(result, ScalarValue::Int64(0)));
}
#[test]
fn aggregate_single_value() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let array = Int64Array::from(vec![42_i64]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let mut state = AggState::default();
udf.accumulate(&mut state, &batch).unwrap();
let result = udf.finalize(state).unwrap();
assert!(matches!(result, ScalarValue::Int64(42)));
}
#[test]
fn aggregate_negative_values() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let array = Int64Array::from(vec![-10_i64, -20, -30]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let mut state = AggState::default();
udf.accumulate(&mut state, &batch).unwrap();
let result = udf.finalize(state).unwrap();
assert!(matches!(result, ScalarValue::Int64(-60)));
}
#[test]
fn aggregate_mixed_positive_negative() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let array = Int64Array::from(vec![-5_i64, 10, -3, 8]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let mut state = AggState::default();
udf.accumulate(&mut state, &batch).unwrap();
let result = udf.finalize(state).unwrap();
assert!(matches!(result, ScalarValue::Int64(10)));
}
#[test]
fn aggregate_multiple_accumulations() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let b1 = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![1_i64, 2]))],
)
.unwrap();
let b2 = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![3_i64, 4]))],
)
.unwrap();
let mut state = AggState::default();
udf.accumulate(&mut state, &b1).unwrap();
udf.accumulate(&mut state, &b2).unwrap();
let result = udf.finalize(state).unwrap();
assert!(matches!(result, ScalarValue::Int64(10)));
}
#[test]
fn aggregate_wrong_type_in_batch() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, true)]));
let array = arrow::array::StringArray::from(vec!["hello"]);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let mut state = AggState::default();
let err = udf.accumulate(&mut state, &batch).unwrap_err();
assert!(matches!(err, UdfError::InvalidArgument { .. }));
}
#[test]
fn aggregate_name_and_schemas() {
let udf = SumAggUdf::new();
assert_eq!(udf.name(), "sum_agg");
assert_eq!(udf.input_schema().fields().len(), 1);
assert_eq!(udf.output_field().name(), "sum");
}
#[test]
fn table_udf_name_and_schema() {
let udf = ConstantTableUdf::new(99);
assert_eq!(udf.name(), "constant_table");
assert_eq!(udf.output_schema().fields().len(), 1);
assert_eq!(udf.output_schema().field(0).name(), "constant");
}
#[test]
fn table_udf_ignores_args() {
let udf = ConstantTableUdf::new(7);
let args = vec![
ScalarValue::Int64(1),
ScalarValue::Utf8("hello".into()),
ScalarValue::Boolean(true),
];
let batch = udf.call(&args).unwrap();
let col = batch
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
assert_eq!(col.value(0), 7);
}
#[test]
fn scalar_value_variants() {
let null = ScalarValue::Null;
let int = ScalarValue::Int64(42);
let float = ScalarValue::Float64(3.15);
let utf8 = ScalarValue::Utf8("hello".into());
let bool = ScalarValue::Boolean(true);
let bytes = ScalarValue::Bytes(vec![1, 2, 3]);
assert!(format!("{:?}", null).contains("Null"));
assert!(format!("{:?}", int).contains("42"));
assert!(format!("{:?}", float).contains("3.15"));
assert!(format!("{:?}", utf8).contains("hello"));
assert!(format!("{:?}", bool).contains("true"));
assert!(format!("{:?}", bytes).contains("Bytes"));
}
#[test]
fn scalar_value_clone() {
let v = ScalarValue::Utf8("test".into());
let c = v.clone();
assert!(matches!(c, ScalarValue::Utf8(s) if s == "test"));
}
#[test]
fn agg_state_default_is_empty() {
let s = AggState::default();
assert!(s.data.is_empty());
}
#[test]
fn agg_state_debug() {
let s = AggState {
data: vec![1, 2, 3],
};
let debug = format!("{:?}", s);
assert!(debug.contains("1, 2, 3"));
}
#[test]
fn udf_error_is_std_error() {
let err: Box<dyn std::error::Error> = Box::new(UdfError::Arrow("test".into()));
assert!(!err.to_string().is_empty());
}
#[test]
fn arrow_error_conversion() {
let arrow_err = arrow::error::ArrowError::InvalidArgumentError("bad".into());
let udf_err: UdfError = arrow_err.into();
assert!(matches!(udf_err, UdfError::Arrow(_)));
assert!(udf_err.to_string().contains("bad"));
}
#[test]
fn multiply_scalar_large_batch() {
let udf = MultiplyScalarUdf::new("big", "x", 7);
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)]));
let values: Vec<i64> = (0..10000).collect();
let array = Int64Array::from(values);
let batch = RecordBatch::try_new(schema, vec![Arc::new(array)]).unwrap();
let result = udf.call(&batch).unwrap();
let arr = result.as_any().downcast_ref::<Int64Array>().unwrap();
assert_eq!(arr.len(), 10000);
assert_eq!(arr.value(0), 0);
assert_eq!(arr.value(1), 7);
assert_eq!(arr.value(9999), 9999 * 7);
}
#[test]
fn registry_new_is_empty() {
let registry = UdfRegistry::new();
assert!(registry.scalar_names().is_empty());
assert!(registry.aggregate_names().is_empty());
assert!(registry.table_names().is_empty());
}
#[test]
fn registry_default_is_empty() {
let registry = UdfRegistry::default();
assert!(registry.scalar_names().is_empty());
}
#[test]
fn aggregate_merge_symmetric() {
let udf = SumAggUdf::new();
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
DataType::Int64,
true,
)]));
let b1 = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![10_i64]))],
)
.unwrap();
let b2 = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(vec![20_i64]))],
)
.unwrap();
let mut s1 = AggState::default();
let mut s2 = AggState::default();
udf.accumulate(&mut s1, &b1).unwrap();
udf.accumulate(&mut s2, &b2).unwrap();
let m12 = udf.merge(s1.clone(), s2.clone()).unwrap();
let m21 = udf.merge(s2, s1).unwrap();
let r12 = udf.finalize(m12).unwrap();
let r21 = udf.finalize(m21).unwrap();
assert!(matches!(r12, ScalarValue::Int64(30)));
assert!(matches!(r21, ScalarValue::Int64(30)));
}
}
#[derive(Clone, Debug, Default)]
pub struct ResourceLimits {
pub max_memory_bytes: Option<u64>,
pub max_execution_time_ms: Option<u64>,
}
pub trait SandboxedUdfExecutor: Send + Sync {
fn execute_with_limits(
&self,
udf: &dyn ScalarUdf,
batch: &RecordBatch,
limits: &ResourceLimits,
) -> Result<ArrayRef, UdfError>;
}
pub struct DefaultSandboxedExecutor;
impl SandboxedUdfExecutor for DefaultSandboxedExecutor {
fn execute_with_limits(
&self,
udf: &dyn ScalarUdf,
batch: &RecordBatch,
limits: &ResourceLimits,
) -> Result<ArrayRef, UdfError> {
if krishiv_common::profile_forbids_native_scalar_udfs(
krishiv_common::resolve_durability_profile(),
) {
return Err(UdfError::Execution {
message: String::from(
"native UDF execution runs with full process privileges; under durable \
profiles use LANGUAGE sql UDFs or set KRISHIV_ALLOW_FULL_PRIVILEGE_UDFS=1",
),
});
}
let start = std::time::Instant::now();
let result =
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| udf.call(batch))) {
Ok(Ok(array)) => array,
Ok(Err(error)) => return Err(error),
Err(payload) => {
let message = krishiv_common::panic_payload_to_string(&*payload);
return Err(UdfError::Panic(format!(
"UDF '{}' panicked during execution: {}",
udf.name(),
message
)));
}
};
if let Some(max_ms) = limits.max_execution_time_ms
&& start.elapsed().as_millis() as u64 > max_ms
{
return Err(UdfError::Execution {
message: format!("UDF exceeded time limit of {} ms", max_ms),
});
}
if let Some(max_bytes) = limits.max_memory_bytes {
let approx_bytes: usize = batch
.columns()
.iter()
.map(|c| c.get_array_memory_size())
.sum();
if approx_bytes as u64 > max_bytes {
return Err(UdfError::Execution {
message: format!(
"UDF input exceeded memory limit of {} bytes (approx {} bytes)",
max_bytes, approx_bytes
),
});
}
let output_size: usize = result.get_array_memory_size();
if output_size as u64 > max_bytes {
return Err(UdfError::Execution {
message: format!(
"UDF output exceeded memory limit of {} bytes (approx {} bytes)",
max_bytes, output_size
),
});
}
}
Ok(result)
}
}
#[cfg(test)]
mod memory_enforcement_tests {
use super::*;
use arrow::array::{ArrayRef, Int64Array};
use arrow::datatypes::{DataType, Field, Schema};
use std::sync::Arc;
#[derive(Debug)]
struct IdentityHeavyUdf {
name: String,
schema: Schema,
}
impl IdentityHeavyUdf {
fn new() -> Self {
let schema = Schema::new(vec![Field::new("a", DataType::Int64, false)]);
Self {
name: "identity_heavy".to_string(),
schema,
}
}
}
impl ScalarUdf for IdentityHeavyUdf {
fn name(&self) -> &str {
&self.name
}
fn input_schema(&self) -> &Schema {
&self.schema
}
fn output_field(&self) -> &Field {
self.schema.field(0)
}
fn call(&self, batch: &RecordBatch) -> Result<ArrayRef, UdfError> {
Ok(batch.column(0).clone())
}
}
#[test]
fn default_sandboxed_executor_enforces_memory_limit() {
let mut registry = UdfRegistry::new();
let udf = Arc::new(IdentityHeavyUdf::new());
registry.register_scalar(udf.clone());
let col = Int64Array::from(vec![1, 2, 3, 4, 5]);
let batch = RecordBatch::try_new(
Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)])),
vec![Arc::new(col)],
)
.unwrap();
let executor = DefaultSandboxedExecutor;
let limits = ResourceLimits {
max_memory_bytes: Some(1),
max_execution_time_ms: None,
};
let err = registry
.execute_scalar_with_limits("identity_heavy", &batch, &limits, &executor)
.unwrap_err();
match err {
UdfError::Execution { message } => {
assert!(
message.contains("exceeded memory limit"),
"expected memory limit error, got: {}",
message
);
}
other => panic!("expected Execution error, got {:?}", other),
}
}
#[derive(Debug)]
struct PanickingUdf;
impl ScalarUdf for PanickingUdf {
fn name(&self) -> &str {
"panicking_udf"
}
fn input_schema(&self) -> &Schema {
static SCHEMA: std::sync::OnceLock<Schema> = std::sync::OnceLock::new();
SCHEMA.get_or_init(|| Schema::new(vec![Field::new("x", DataType::Int64, true)]))
}
fn output_field(&self) -> &Field {
static FIELD: std::sync::OnceLock<Field> = std::sync::OnceLock::new();
FIELD.get_or_init(|| Field::new("x", DataType::Int64, true))
}
fn call(&self, _batch: &RecordBatch) -> Result<ArrayRef, UdfError> {
panic!("deliberate test panic: kaboom");
}
}
#[test]
fn default_sandboxed_executor_catches_udf_panic() {
let mut registry = UdfRegistry::new();
registry.register_scalar(Arc::new(PanickingUdf));
let executor = DefaultSandboxedExecutor;
let limits = ResourceLimits::default();
let batch = RecordBatch::try_new(
Arc::new(Schema::new(vec![Field::new("x", DataType::Int64, true)])),
vec![Arc::new(Int64Array::from(vec![1, 2, 3]))],
)
.unwrap();
let err = registry
.execute_scalar_with_limits("panicking_udf", &batch, &limits, &executor)
.unwrap_err();
match err {
UdfError::Panic(message) => {
assert!(
message.contains("panicking_udf"),
"udf name in error: {message}"
);
assert!(
message.contains("kaboom"),
"panic message in error: {message}"
);
}
other => panic!("expected Panic error, got {other:?}"),
}
}
#[test]
fn panic_message_extracts_str_payload() {
let payload: Box<dyn std::any::Any + Send> = Box::new("static str payload");
assert_eq!(
krishiv_common::panic_payload_to_string(&*payload),
"static str payload"
);
}
#[test]
fn panic_message_extracts_string_payload() {
let payload: Box<dyn std::any::Any + Send> = Box::new(String::from("owned payload"));
assert_eq!(
krishiv_common::panic_payload_to_string(&*payload),
"owned payload"
);
}
#[test]
fn panic_message_falls_back_for_unknown_payloads() {
let payload: Box<dyn std::any::Any + Send> = Box::new(42u32);
assert_eq!(
krishiv_common::panic_payload_to_string(&*payload),
"non-string panic payload"
);
}
}