use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use super::sort::{compare_values, Compare};
use super::Value;
pub type AggregateFn = Arc<dyn Fn(&[&Value]) -> Value + Send + Sync>;
#[derive(Clone)]
enum Kind {
Count,
Sum,
Min,
Max,
Mean,
Custom(AggregateFn),
}
#[derive(Clone)]
pub struct Aggregate {
column: usize,
kind: Kind,
}
impl fmt::Debug for Aggregate {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let kind = match self.kind {
Kind::Count => "count",
Kind::Sum => "sum",
Kind::Min => "min",
Kind::Max => "max",
Kind::Mean => "mean",
Kind::Custom(_) => "custom",
};
f.debug_struct("Aggregate")
.field("column", &self.column)
.field("kind", &kind)
.finish()
}
}
impl Aggregate {
fn new(column: usize, kind: Kind) -> Self {
Aggregate { column, kind }
}
pub fn count(column: usize) -> Self {
Self::new(column, Kind::Count)
}
pub fn sum(column: usize) -> Self {
Self::new(column, Kind::Sum)
}
pub fn min(column: usize) -> Self {
Self::new(column, Kind::Min)
}
pub fn max(column: usize) -> Self {
Self::new(column, Kind::Max)
}
pub fn mean(column: usize) -> Self {
Self::new(column, Kind::Mean)
}
pub fn custom(column: usize, f: impl Fn(&[&Value]) -> Value + Send + Sync + 'static) -> Self {
Self::new(column, Kind::Custom(Arc::new(f)))
}
pub fn column(&self) -> usize {
self.column
}
pub fn is_count(&self) -> bool {
matches!(self.kind, Kind::Count)
}
pub fn compute(&self, values: &[&Value]) -> Value {
let non_empty = || values.iter().copied().filter(|v| !v.is_empty());
match &self.kind {
Kind::Count => Value::from(non_empty().count()),
Kind::Sum => {
let numbers: Vec<&Value> = values
.iter()
.copied()
.filter(|v| v.as_f64().is_some())
.collect();
if numbers.is_empty() {
return Value::Null;
}
let ints: Option<i64> = numbers.iter().try_fold(0i64, |sum, v| match v {
Value::Int(n) => sum.checked_add(*n),
_ => None,
});
ints.map_or_else(
|| Value::Float(numbers.iter().filter_map(|v| v.as_f64()).sum()),
Value::Int,
)
}
Kind::Min => non_empty()
.min_by(|a, b| compare_values(a, b, Compare::Natural))
.cloned()
.unwrap_or_default(),
Kind::Max => non_empty()
.reduce(|best, v| {
if compare_values(v, best, Compare::Natural).is_gt() {
v
} else {
best
}
})
.cloned()
.unwrap_or_default(),
Kind::Mean => {
let numbers: Vec<f64> = values.iter().filter_map(|v| v.as_f64()).collect();
if numbers.is_empty() {
Value::Null
} else {
Value::Float(numbers.iter().sum::<f64>() / numbers.len() as f64)
}
}
Kind::Custom(f) => f(values),
}
}
pub fn over<R: AsRef<[Value]>>(&self, data: &[R], rows: &[usize]) -> Value {
const NULL: Value = Value::Null;
let cells: Vec<&Value> = rows
.iter()
.map(|&i| data[i].as_ref().get(self.column).unwrap_or(&NULL))
.collect();
self.compute(&cells)
}
}
#[derive(Clone, Debug)]
pub struct GroupBy {
column: usize,
aggregates: Vec<Aggregate>,
label: String,
}
impl GroupBy {
pub fn new(column: usize) -> Self {
GroupBy {
column,
aggregates: Vec::new(),
label: "subtotal".to_string(),
}
}
pub fn aggregate(mut self, aggregate: Aggregate) -> Self {
self.aggregates.push(aggregate);
self
}
pub fn label(mut self, label: impl Into<String>) -> Self {
self.label = label.into();
self
}
pub fn column(&self) -> usize {
self.column
}
pub fn aggregates(&self) -> &[Aggregate] {
&self.aggregates
}
pub fn summary_label(&self) -> &str {
&self.label
}
pub fn groups<R: AsRef<[Value]>>(&self, rows: &[R], order: &[usize]) -> Vec<Group> {
let mut groups: Vec<Group> = Vec::new();
let mut by_key: HashMap<String, usize> = HashMap::new();
for &row in order {
let key = rows[row]
.as_ref()
.get(self.column)
.cloned()
.unwrap_or_default();
let slot = *by_key.entry(key.plain()).or_insert_with(|| {
groups.push(Group {
key,
rows: Vec::new(),
aggregates: Vec::new(),
});
groups.len() - 1
});
groups[slot].rows.push(row);
}
for group in &mut groups {
group.aggregates = self
.aggregates
.iter()
.map(|aggregate| aggregate.over(rows, &group.rows))
.collect();
}
groups
}
}
#[derive(Clone, Debug)]
pub struct Group {
pub key: Value,
pub rows: Vec<usize>,
pub aggregates: Vec<Value>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn aggregates_over_mixed_cells() {
let cells = [
Value::Int(3),
Value::Null,
Value::Float(1.5),
Value::from("x"),
Value::Int(-2),
];
let refs: Vec<&Value> = cells.iter().collect();
assert_eq!(Aggregate::count(0).compute(&refs), Value::Int(4));
assert_eq!(Aggregate::sum(0).compute(&refs), Value::Float(2.5));
assert_eq!(Aggregate::min(0).compute(&refs), Value::Int(-2));
assert_eq!(Aggregate::max(0).compute(&refs), Value::from("x"));
assert_eq!(Aggregate::mean(0).compute(&refs), Value::Float(2.5 / 3.0));
assert_eq!(Aggregate::sum(0).compute(&[]), Value::Null);
let ints = [Value::Int(i64::MAX), Value::Int(1)];
let refs: Vec<&Value> = ints.iter().collect();
assert_eq!(
Aggregate::sum(0).compute(&refs),
Value::Float(i64::MAX as f64 + 1.0)
);
}
}