use anyhow::Result;
use crate::exec::function::{Accumulator, AggregateFunction, Signature};
use crate::expr::Kind;
use crate::val::{Number, Value};
#[derive(Debug, Clone, Copy, Default)]
pub struct Count;
impl AggregateFunction for Count {
fn name(&self) -> &'static str {
"count"
}
fn create_accumulator(&self) -> Box<dyn Accumulator> {
Box::new(CountAccumulator::default())
}
fn signature(&self) -> Signature {
Signature::new().returns(Kind::Int)
}
}
#[derive(Debug, Clone, Default)]
struct CountAccumulator {
count: i64,
}
impl Accumulator for CountAccumulator {
fn update(&mut self, _value: Value) -> Result<()> {
self.count += 1;
Ok(())
}
fn update_batch(&mut self, values: &[Value]) -> Result<()> {
self.count += values.len() as i64;
Ok(())
}
fn merge(&mut self, other: Box<dyn Accumulator>) -> Result<()> {
let other = other
.as_any()
.downcast_ref::<CountAccumulator>()
.ok_or_else(|| anyhow::anyhow!("Cannot merge incompatible accumulators"))?;
self.count += other.count;
Ok(())
}
fn finalize(&self) -> Result<Value> {
Ok(Value::Number(Number::Int(self.count)))
}
fn reset(&mut self) {
self.count = 0;
}
fn clone_box(&self) -> Box<dyn Accumulator> {
Box::new(self.clone())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct CountField;
impl AggregateFunction for CountField {
fn name(&self) -> &'static str {
"count"
}
fn create_accumulator(&self) -> Box<dyn Accumulator> {
Box::new(CountFieldAccumulator::default())
}
fn signature(&self) -> Signature {
Signature::new().arg("value", Kind::Any).returns(Kind::Int)
}
}
#[derive(Debug, Clone, Default)]
struct CountFieldAccumulator {
count: i64,
}
impl Accumulator for CountFieldAccumulator {
fn update(&mut self, value: Value) -> Result<()> {
if value.is_truthy() {
self.count += 1;
}
Ok(())
}
fn update_batch(&mut self, values: &[Value]) -> Result<()> {
self.count += values.iter().filter(|v| v.is_truthy()).count() as i64;
Ok(())
}
fn merge(&mut self, other: Box<dyn Accumulator>) -> Result<()> {
let other = other
.as_any()
.downcast_ref::<CountFieldAccumulator>()
.ok_or_else(|| anyhow::anyhow!("Cannot merge incompatible accumulators"))?;
self.count += other.count;
Ok(())
}
fn finalize(&self) -> Result<Value> {
Ok(Value::Number(Number::Int(self.count)))
}
fn reset(&mut self) {
self.count = 0;
}
fn clone_box(&self) -> Box<dyn Accumulator> {
Box::new(self.clone())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[cfg(test)]
mod tests {
use surrealdb_strand::Strand;
use super::*;
fn as_int(v: &Value) -> i64 {
match v {
Value::Number(Number::Int(i)) => *i,
_ => panic!("Expected Int, got {:?}", v),
}
}
#[test]
fn count_zero_items() {
let func = Count;
let acc = func.create_accumulator();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 0);
}
#[test]
fn count_single_item() {
let func = Count;
let mut acc = func.create_accumulator();
acc.update(Value::Number(Number::Int(42))).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 1);
}
#[test]
fn count_multiple_items() {
let func = Count;
let mut acc = func.create_accumulator();
acc.update(Value::Number(Number::Int(1))).unwrap();
acc.update(Value::Number(Number::Int(2))).unwrap();
acc.update(Value::Number(Number::Int(3))).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 3);
}
#[test]
fn count_merge() {
let func = Count;
let mut acc1 = func.create_accumulator();
acc1.update(Value::Number(Number::Int(1))).unwrap();
acc1.update(Value::Number(Number::Int(2))).unwrap();
let mut acc2 = func.create_accumulator();
acc2.update(Value::Number(Number::Int(3))).unwrap();
acc1.merge(acc2).unwrap();
let result = acc1.finalize().unwrap();
assert_eq!(as_int(&result), 3);
}
#[test]
fn count_reset() {
let func = Count;
let mut acc = func.create_accumulator();
acc.update(Value::Number(Number::Int(1))).unwrap();
acc.update(Value::Number(Number::Int(2))).unwrap();
acc.reset();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 0);
}
#[test]
fn count_field_zero_items() {
let func = CountField;
let acc = func.create_accumulator();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 0);
}
#[test]
fn count_field_truthy_value() {
let func = CountField;
let mut acc = func.create_accumulator();
acc.update(Value::Number(Number::Int(42))).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 1);
}
#[test]
fn count_field_falsy_values() {
let func = CountField;
let mut acc = func.create_accumulator();
acc.update(Value::None).unwrap();
acc.update(Value::Bool(false)).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 0);
}
#[test]
fn count_field_mixed_values() {
let func = CountField;
let mut acc = func.create_accumulator();
acc.update(Value::Number(Number::Int(1))).unwrap(); acc.update(Value::None).unwrap(); acc.update(Value::String(Strand::new_static("hello"))).unwrap(); acc.update(Value::Bool(false)).unwrap(); acc.update(Value::Bool(true)).unwrap(); let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 3);
}
#[test]
fn count_field_merge() {
let func = CountField;
let mut acc1 = func.create_accumulator();
acc1.update(Value::Number(Number::Int(1))).unwrap();
acc1.update(Value::None).unwrap();
let mut acc2 = func.create_accumulator();
acc2.update(Value::String(Strand::new_static("test"))).unwrap();
acc1.merge(acc2).unwrap();
let result = acc1.finalize().unwrap();
assert_eq!(as_int(&result), 2);
}
#[test]
fn count_batch_empty() {
let func = Count;
let mut acc = func.create_accumulator();
acc.update_batch(&[]).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 0);
}
#[test]
fn count_batch_multiple() {
let func = Count;
let mut acc = func.create_accumulator();
let values = vec![
Value::Number(Number::Int(1)),
Value::None,
Value::String(Strand::new_static("test")),
];
acc.update_batch(&values).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 3);
}
#[test]
fn count_batch_then_single() {
let func = Count;
let mut acc = func.create_accumulator();
let values = vec![Value::Number(Number::Int(1)), Value::Number(Number::Int(2))];
acc.update_batch(&values).unwrap();
acc.update(Value::Number(Number::Int(3))).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 3);
}
#[test]
fn count_field_batch_empty() {
let func = CountField;
let mut acc = func.create_accumulator();
acc.update_batch(&[]).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 0);
}
#[test]
fn count_field_batch_mixed() {
let func = CountField;
let mut acc = func.create_accumulator();
let values = vec![
Value::Number(Number::Int(1)), Value::None, Value::String(Strand::new_static("hello")), Value::Bool(false), Value::Bool(true), ];
acc.update_batch(&values).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 3);
}
#[test]
fn count_field_batch_then_single() {
let func = CountField;
let mut acc = func.create_accumulator();
let values = vec![Value::Number(Number::Int(1)), Value::None];
acc.update_batch(&values).unwrap();
acc.update(Value::Bool(true)).unwrap();
let result = acc.finalize().unwrap();
assert_eq!(as_int(&result), 2);
}
}