1use std::collections::HashMap;
27use std::fmt;
28use std::sync::Arc;
29
30use super::sort::{compare_values, Compare};
31use super::Value;
32
33pub type AggregateFn = Arc<dyn Fn(&[&Value]) -> Value + Send + Sync>;
35
36#[derive(Clone)]
37enum Kind {
38 Count,
39 Sum,
40 Min,
41 Max,
42 Mean,
43 Custom(AggregateFn),
44}
45
46#[derive(Clone)]
48pub struct Aggregate {
49 column: usize,
50 kind: Kind,
51}
52
53impl fmt::Debug for Aggregate {
54 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
55 let kind = match self.kind {
56 Kind::Count => "count",
57 Kind::Sum => "sum",
58 Kind::Min => "min",
59 Kind::Max => "max",
60 Kind::Mean => "mean",
61 Kind::Custom(_) => "custom",
62 };
63 f.debug_struct("Aggregate")
64 .field("column", &self.column)
65 .field("kind", &kind)
66 .finish()
67 }
68}
69
70impl Aggregate {
71 fn new(column: usize, kind: Kind) -> Self {
72 Aggregate { column, kind }
73 }
74
75 pub fn count(column: usize) -> Self {
78 Self::new(column, Kind::Count)
79 }
80
81 pub fn sum(column: usize) -> Self {
84 Self::new(column, Kind::Sum)
85 }
86
87 pub fn min(column: usize) -> Self {
89 Self::new(column, Kind::Min)
90 }
91
92 pub fn max(column: usize) -> Self {
94 Self::new(column, Kind::Max)
95 }
96
97 pub fn mean(column: usize) -> Self {
99 Self::new(column, Kind::Mean)
100 }
101
102 pub fn custom(column: usize, f: impl Fn(&[&Value]) -> Value + Send + Sync + 'static) -> Self {
104 Self::new(column, Kind::Custom(Arc::new(f)))
105 }
106
107 pub fn column(&self) -> usize {
109 self.column
110 }
111
112 pub fn is_count(&self) -> bool {
115 matches!(self.kind, Kind::Count)
116 }
117
118 pub fn compute(&self, values: &[&Value]) -> Value {
120 let non_empty = || values.iter().copied().filter(|v| !v.is_empty());
121 match &self.kind {
122 Kind::Count => Value::from(non_empty().count()),
123 Kind::Sum => {
124 let numbers: Vec<&Value> = values
125 .iter()
126 .copied()
127 .filter(|v| v.as_f64().is_some())
128 .collect();
129 if numbers.is_empty() {
130 return Value::Null;
131 }
132 let ints: Option<i64> = numbers.iter().try_fold(0i64, |sum, v| match v {
133 Value::Int(n) => sum.checked_add(*n),
134 _ => None,
135 });
136 ints.map_or_else(
137 || Value::Float(numbers.iter().filter_map(|v| v.as_f64()).sum()),
138 Value::Int,
139 )
140 }
141 Kind::Min => non_empty()
142 .min_by(|a, b| compare_values(a, b, Compare::Natural))
143 .cloned()
144 .unwrap_or_default(),
145 Kind::Max => non_empty()
146 .reduce(|best, v| {
148 if compare_values(v, best, Compare::Natural).is_gt() {
149 v
150 } else {
151 best
152 }
153 })
154 .cloned()
155 .unwrap_or_default(),
156 Kind::Mean => {
157 let numbers: Vec<f64> = values.iter().filter_map(|v| v.as_f64()).collect();
158 if numbers.is_empty() {
159 Value::Null
160 } else {
161 Value::Float(numbers.iter().sum::<f64>() / numbers.len() as f64)
162 }
163 }
164 Kind::Custom(f) => f(values),
165 }
166 }
167
168 pub fn over<R: AsRef<[Value]>>(&self, data: &[R], rows: &[usize]) -> Value {
170 const NULL: Value = Value::Null;
171 let cells: Vec<&Value> = rows
172 .iter()
173 .map(|&i| data[i].as_ref().get(self.column).unwrap_or(&NULL))
174 .collect();
175 self.compute(&cells)
176 }
177}
178
179#[derive(Clone, Debug)]
181pub struct GroupBy {
182 column: usize,
183 aggregates: Vec<Aggregate>,
184 label: String,
185}
186
187impl GroupBy {
188 pub fn new(column: usize) -> Self {
190 GroupBy {
191 column,
192 aggregates: Vec::new(),
193 label: "subtotal".to_string(),
194 }
195 }
196
197 pub fn aggregate(mut self, aggregate: Aggregate) -> Self {
199 self.aggregates.push(aggregate);
200 self
201 }
202
203 pub fn label(mut self, label: impl Into<String>) -> Self {
206 self.label = label.into();
207 self
208 }
209
210 pub fn column(&self) -> usize {
212 self.column
213 }
214
215 pub fn aggregates(&self) -> &[Aggregate] {
217 &self.aggregates
218 }
219
220 pub fn summary_label(&self) -> &str {
222 &self.label
223 }
224
225 pub fn groups<R: AsRef<[Value]>>(&self, rows: &[R], order: &[usize]) -> Vec<Group> {
230 let mut groups: Vec<Group> = Vec::new();
231 let mut by_key: HashMap<String, usize> = HashMap::new();
232 for &row in order {
233 let key = rows[row]
234 .as_ref()
235 .get(self.column)
236 .cloned()
237 .unwrap_or_default();
238 let slot = *by_key.entry(key.plain()).or_insert_with(|| {
239 groups.push(Group {
240 key,
241 rows: Vec::new(),
242 aggregates: Vec::new(),
243 });
244 groups.len() - 1
245 });
246 groups[slot].rows.push(row);
247 }
248 for group in &mut groups {
249 group.aggregates = self
250 .aggregates
251 .iter()
252 .map(|aggregate| aggregate.over(rows, &group.rows))
253 .collect();
254 }
255 groups
256 }
257}
258
259#[derive(Clone, Debug)]
261pub struct Group {
262 pub key: Value,
264 pub rows: Vec<usize>,
266 pub aggregates: Vec<Value>,
268}
269
270#[cfg(test)]
271mod tests {
272 use super::*;
273
274 #[test]
275 fn aggregates_over_mixed_cells() {
276 let cells = [
277 Value::Int(3),
278 Value::Null,
279 Value::Float(1.5),
280 Value::from("x"),
281 Value::Int(-2),
282 ];
283 let refs: Vec<&Value> = cells.iter().collect();
284 assert_eq!(Aggregate::count(0).compute(&refs), Value::Int(4));
285 assert_eq!(Aggregate::sum(0).compute(&refs), Value::Float(2.5));
286 assert_eq!(Aggregate::min(0).compute(&refs), Value::Int(-2));
287 assert_eq!(Aggregate::max(0).compute(&refs), Value::from("x"));
288 assert_eq!(Aggregate::mean(0).compute(&refs), Value::Float(2.5 / 3.0));
289 assert_eq!(Aggregate::sum(0).compute(&[]), Value::Null);
290 let ints = [Value::Int(i64::MAX), Value::Int(1)];
291 let refs: Vec<&Value> = ints.iter().collect();
292 assert_eq!(
293 Aggregate::sum(0).compute(&refs),
294 Value::Float(i64::MAX as f64 + 1.0)
295 );
296 }
297}