use crate::error::SqawkResult;
use crate::table::Value;
#[derive(Debug, Clone, Copy)]
pub enum AggregateFunction {
Count,
Sum,
Avg,
Min,
Max,
}
impl AggregateFunction {
pub fn from_name(name: &str) -> Option<Self> {
match name.to_uppercase().as_str() {
"COUNT" => Some(AggregateFunction::Count),
"SUM" => Some(AggregateFunction::Sum),
"AVG" => Some(AggregateFunction::Avg),
"MIN" => Some(AggregateFunction::Min),
"MAX" => Some(AggregateFunction::Max),
_ => None,
}
}
pub fn execute(&self, values: &[Value]) -> SqawkResult<Value> {
match self {
AggregateFunction::Count => self.count(values),
AggregateFunction::Sum => self.sum(values),
AggregateFunction::Avg => self.avg(values),
AggregateFunction::Min => self.min(values),
AggregateFunction::Max => self.max(values),
}
}
fn count(&self, values: &[Value]) -> SqawkResult<Value> {
let count = values.iter().filter(|v| !matches!(v, Value::Null)).count();
Ok(Value::Integer(count as i64))
}
fn sum(&self, values: &[Value]) -> SqawkResult<Value> {
let mut is_float = false;
let mut int_sum: i64 = 0;
let mut float_sum: f64 = 0.0;
let mut count = 0;
for value in values {
match value {
Value::Integer(i) => {
if is_float {
float_sum += *i as f64;
} else {
int_sum += *i;
}
count += 1;
}
Value::Float(f) => {
if !is_float {
float_sum = int_sum as f64;
is_float = true;
}
float_sum += *f;
count += 1;
}
_ => {}
}
}
if count == 0 {
return Ok(Value::Null);
}
if is_float {
Ok(Value::Float(float_sum))
} else {
Ok(Value::Integer(int_sum))
}
}
fn avg(&self, values: &[Value]) -> SqawkResult<Value> {
let sum = self.sum(values)?;
let count = values
.iter()
.filter(|v| matches!(v, Value::Integer(_) | Value::Float(_)))
.count();
if count == 0 {
return Ok(Value::Null);
}
match sum {
Value::Integer(i) => Ok(Value::Float(i as f64 / count as f64)),
Value::Float(f) => Ok(Value::Float(f / count as f64)),
_ => Ok(Value::Null), }
}
fn min(&self, values: &[Value]) -> SqawkResult<Value> {
let non_null_values: Vec<&Value> = values
.iter()
.filter(|v| !matches!(v, Value::Null))
.collect();
if non_null_values.is_empty() {
return Ok(Value::Null);
}
let mut min_value = non_null_values[0].clone();
for value in &non_null_values[1..] {
if value < &&min_value {
min_value = (*value).clone();
}
}
Ok(min_value)
}
fn max(&self, values: &[Value]) -> SqawkResult<Value> {
let non_null_values: Vec<&Value> = values
.iter()
.filter(|v| !matches!(v, Value::Null))
.collect();
if non_null_values.is_empty() {
return Ok(Value::Null);
}
let mut max_value = non_null_values[0].clone();
for value in &non_null_values[1..] {
if value > &&max_value {
max_value = (*value).clone();
}
}
Ok(max_value)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_count_function() {
let values = vec![
Value::Integer(10),
Value::Null,
Value::Integer(20),
Value::String("test".to_string().into()),
];
let count = AggregateFunction::Count.execute(&values).unwrap();
assert_eq!(count, Value::Integer(3));
}
#[test]
fn test_sum_function() {
let values = vec![
Value::Integer(10),
Value::Null,
Value::Integer(20),
Value::Float(5.5),
];
let sum = AggregateFunction::Sum.execute(&values).unwrap();
assert_eq!(sum, Value::Float(35.5));
}
#[test]
fn test_sum_integers_only() {
let values = vec![Value::Integer(10), Value::Integer(20), Value::Integer(30)];
let sum = AggregateFunction::Sum.execute(&values).unwrap();
assert_eq!(sum, Value::Integer(60));
}
#[test]
fn test_avg_function() {
let values = vec![
Value::Integer(10),
Value::Null,
Value::Integer(20),
Value::Float(30.0),
];
let avg = AggregateFunction::Avg.execute(&values).unwrap();
if let Value::Float(f) = avg {
assert!((f - 20.0).abs() < f64::EPSILON);
} else {
panic!("Expected Float, got {:?}", avg);
}
}
#[test]
fn test_min_function() {
let values = vec![
Value::Integer(30),
Value::Integer(10),
Value::Null,
Value::Integer(20),
];
let min = AggregateFunction::Min.execute(&values).unwrap();
assert_eq!(min, Value::Integer(10));
}
#[test]
fn test_min_with_mixed_types() {
let values = vec![
Value::Integer(30),
Value::Float(5.5),
Value::String("abc".to_string().into()),
Value::Null,
];
let min = AggregateFunction::Min.execute(&values).unwrap();
assert_eq!(min, Value::Float(5.5));
}
#[test]
fn test_max_function() {
let values = vec![
Value::Integer(30),
Value::Integer(10),
Value::Null,
Value::Integer(20),
];
let max = AggregateFunction::Max.execute(&values).unwrap();
assert_eq!(max, Value::Integer(30));
}
#[test]
fn test_max_with_mixed_types() {
let values = vec![
Value::Integer(30),
Value::Float(50.5),
Value::String("xyz".to_string().into()),
Value::Null,
];
let max = AggregateFunction::Max.execute(&values).unwrap();
assert_eq!(max, Value::String("xyz".to_string().into()));
}
}