1use crate::{TableRule, ValidationError, error::DataFusionSnafu};
2
3use datafusion::functions_aggregate::{count::count_all, expr_fn::*};
4use datafusion::logical_expr::{SortExpr, Subquery};
5use datafusion::prelude::*;
6use snafu::ResultExt;
7use std::sync::Arc;
8
9#[derive(Debug, Clone, Default)]
11pub struct NullCountRule {
12 negated: Option<bool>,
13}
14
15impl TableRule for NullCountRule {
16 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
17 let new_column_name = self.new_column_name(column_name);
18 let subquery = if !self.negated.unwrap_or(false) {
19 df.clone()
20 .aggregate(
21 vec![],
22 vec![
23 count_all().alias("count_all"),
24 count(col(column_name)).alias(new_column_name.as_str()),
25 ],
26 )?
27 .select(vec![
28 col("count_all")
29 .sub(col(new_column_name.as_str()))
30 .alias(new_column_name.as_str()),
31 ])?
32 } else {
33 df.clone()
34 .aggregate(vec![], vec![count(col(column_name)).alias("count_all")])?
35 .select(vec![col("count_all").alias(new_column_name.as_str())])?
36 };
37
38 let subquery_expr = Expr::ScalarSubquery(Subquery {
39 subquery: Arc::new(subquery.logical_plan().clone()),
40 outer_ref_columns: vec![],
41 });
42
43 df.with_column(&new_column_name, subquery_expr)
44 .context(DataFusionSnafu)
45 }
46
47 fn name(&self) -> &str {
48 if self.negated.unwrap_or(false) {
49 "not_null_count"
50 } else {
51 "null_count"
52 }
53 }
54
55 fn new_column_name(&self, column_name: &str) -> String {
56 format!("{}_{}", column_name, self.name())
57 }
58
59 fn description(&self) -> &str {
60 "Counts the number of (not) null values in a column across the entire table"
61 }
62}
63
64pub fn dfq_null_count() -> Arc<NullCountRule> {
65 std::sync::Arc::new(NullCountRule { negated: None })
66}
67
68pub fn dfq_not_null_count() -> Arc<NullCountRule> {
69 std::sync::Arc::new(NullCountRule {
70 negated: Some(true),
71 })
72}
73
74#[derive(Debug, strum::Display, Clone)]
75pub enum CalculationType {
76 Count,
77 CountDistinct,
78 Avg,
79 StdDev,
80 Max,
81 Min,
82 Sum,
83 Median,
84 CovarPop { x: Option<Expr>, y: Option<Expr> },
85 CovarSamp { x: Option<Expr>, y: Option<Expr> },
86 FirstValue(Option<Vec<SortExpr>>),
87 LastValue,
88 NthValue(i64, Option<Vec<SortExpr>>),
89 RegrAvgX { x: Option<Expr>, y: Option<Expr> },
90 RegrAvgY { x: Option<Expr>, y: Option<Expr> },
91 RegrCount { x: Option<Expr>, y: Option<Expr> },
92 RegrIntercept { x: Option<Expr>, y: Option<Expr> },
93 RegrR2 { x: Option<Expr>, y: Option<Expr> },
94 RegrSlope { x: Option<Expr>, y: Option<Expr> },
95 RegrSxx { x: Option<Expr>, y: Option<Expr> },
96 RegrSxy { x: Option<Expr>, y: Option<Expr> },
97 RegrSyy { x: Option<Expr>, y: Option<Expr> },
98 StddevPop,
99 VarPop,
100 VarSamp,
101}
102
103#[derive(Debug, Clone)]
104pub struct CalculationRule {
105 calculation_type: CalculationType,
106}
107
108impl TableRule for CalculationRule {
109 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
110 let new_column_name = self.new_column_name(column_name);
111 let source_column = col(column_name);
112 let calc_expr = match self.calculation_type.clone() {
113 CalculationType::Count => count(source_column),
114 CalculationType::CountDistinct => count_distinct(source_column),
115 CalculationType::Avg => avg(source_column),
116 CalculationType::StdDev => stddev(source_column),
117 CalculationType::Max => max(source_column),
118 CalculationType::Min => min(source_column),
119 CalculationType::Sum => sum(source_column),
120 CalculationType::Median => median(source_column),
121 CalculationType::CovarPop {
122 x: Some(x),
123 y: None,
124 } => covar_pop(source_column, x),
125 CalculationType::CovarPop {
126 x: None,
127 y: Some(y),
128 } => covar_pop(y, source_column),
129 CalculationType::CovarPop { .. } => {
130 return Err(ValidationError::Configuration {
131 message: "CovarPop must have either x or y".to_string(),
132 });
133 }
134 CalculationType::CovarSamp {
135 x: Some(x),
136 y: None,
137 } => covar_samp(source_column, x),
138 CalculationType::CovarSamp {
139 x: None,
140 y: Some(y),
141 } => covar_samp(y, source_column),
142 CalculationType::CovarSamp { .. } => {
143 return Err(ValidationError::Configuration {
144 message: "CovarSamp must have either x or y".to_string(),
145 });
146 }
147 CalculationType::FirstValue(sort_exprs) => first_value(source_column, sort_exprs),
148 CalculationType::LastValue => last_value(vec![source_column]),
149 CalculationType::NthValue(n, sort_exprs) => {
150 nth_value(source_column, n, sort_exprs.unwrap_or_default())
151 }
152 CalculationType::RegrAvgX {
153 x: Some(x),
154 y: None,
155 } => regr_avgx(source_column, x),
156 CalculationType::RegrAvgX {
157 x: None,
158 y: Some(y),
159 } => regr_avgx(y, source_column),
160 CalculationType::RegrAvgX { .. } => {
161 return Err(ValidationError::Configuration {
162 message: "RegrAvgX must have either x or y".to_string(),
163 });
164 }
165 CalculationType::RegrAvgY {
166 x: Some(x),
167 y: None,
168 } => regr_avgy(source_column, x),
169 CalculationType::RegrAvgY {
170 x: None,
171 y: Some(y),
172 } => regr_avgy(y, source_column),
173 CalculationType::RegrAvgY { .. } => {
174 return Err(ValidationError::Configuration {
175 message: "RegrAvgY must have either x or y".to_string(),
176 });
177 }
178 CalculationType::RegrCount {
179 x: Some(x),
180 y: None,
181 } => regr_count(source_column, x),
182 CalculationType::RegrCount {
183 x: None,
184 y: Some(y),
185 } => regr_count(y, source_column),
186 CalculationType::RegrCount { .. } => {
187 return Err(ValidationError::Configuration {
188 message: "RegrCount must have either x or y".to_string(),
189 });
190 }
191 CalculationType::RegrIntercept {
192 x: Some(x),
193 y: None,
194 } => regr_intercept(source_column, x),
195 CalculationType::RegrIntercept {
196 x: None,
197 y: Some(y),
198 } => regr_intercept(y, source_column),
199 CalculationType::RegrIntercept { .. } => {
200 return Err(ValidationError::Configuration {
201 message: "RegrIntercept must have either x or y".to_string(),
202 });
203 }
204 CalculationType::RegrR2 {
205 x: Some(x),
206 y: None,
207 } => regr_r2(source_column, x),
208 CalculationType::RegrR2 {
209 x: None,
210 y: Some(y),
211 } => regr_r2(y, source_column),
212 CalculationType::RegrR2 { .. } => {
213 return Err(ValidationError::Configuration {
214 message: "RegrR2 must have either x or y".to_string(),
215 });
216 }
217 CalculationType::RegrSlope {
218 x: Some(x),
219 y: None,
220 } => regr_slope(source_column, x),
221 CalculationType::RegrSlope {
222 x: None,
223 y: Some(y),
224 } => regr_slope(y, source_column),
225 CalculationType::RegrSlope { .. } => {
226 return Err(ValidationError::Configuration {
227 message: "RegrSlope must have either x or y".to_string(),
228 });
229 }
230 CalculationType::RegrSxx {
231 x: Some(x),
232 y: None,
233 } => regr_sxx(source_column, x),
234 CalculationType::RegrSxx {
235 x: None,
236 y: Some(y),
237 } => regr_sxx(y, source_column),
238 CalculationType::RegrSxx { .. } => {
239 return Err(ValidationError::Configuration {
240 message: "RegrSxx must have either x or y".to_string(),
241 });
242 }
243 CalculationType::RegrSxy {
244 x: Some(x),
245 y: None,
246 } => regr_sxy(source_column, x),
247 CalculationType::RegrSxy {
248 x: None,
249 y: Some(y),
250 } => regr_sxy(y, source_column),
251 CalculationType::RegrSxy { .. } => {
252 return Err(ValidationError::Configuration {
253 message: "RegrSxy must have either x or y".to_string(),
254 });
255 }
256 CalculationType::RegrSyy {
257 x: Some(x),
258 y: None,
259 } => regr_syy(source_column, x),
260 CalculationType::RegrSyy {
261 x: None,
262 y: Some(y),
263 } => regr_syy(y, source_column),
264 CalculationType::RegrSyy { .. } => {
265 return Err(ValidationError::Configuration {
266 message: "RegrSyy must have either x or y".to_string(),
267 });
268 }
269 CalculationType::StddevPop => stddev_pop(source_column),
270 CalculationType::VarPop => var_pop(source_column),
271 CalculationType::VarSamp => var_sample(source_column),
272 }
273 .alias(new_column_name.clone());
274
275 let subq_df = df
276 .clone()
277 .aggregate(vec![], vec![calc_expr])?
278 .select_columns(&[new_column_name.as_str()])?;
279
280 let subq_expr = Expr::ScalarSubquery(Subquery {
281 subquery: Arc::new(subq_df.logical_plan().clone()),
282 outer_ref_columns: vec![],
283 });
284
285 df.with_column(&new_column_name, subq_expr)
286 .context(DataFusionSnafu)
287 }
288
289 fn name(&self) -> &str {
290 "calculation"
291 }
292
293 fn new_column_name(&self, column_name: &str) -> String {
294 format!(
295 "{}_{}",
296 column_name,
297 self.calculation_type.to_string().to_ascii_lowercase()
298 )
299 }
300
301 fn description(&self) -> &str {
302 "Calculates a value for a column across the entire table"
303 }
304}
305
306#[macro_export]
308macro_rules! calc_empty_variant {
309 ($name:ident, $ctype:ident, $(#[$($attrss:tt)*])*) => {
310 $(#[$($attrss)*])*
311 pub fn $name() -> Arc<CalculationRule> {
312 std::sync::Arc::new(CalculationRule { calculation_type: CalculationType::$ctype})
313 }
314 }
315}
316
317#[macro_export]
318macro_rules! calc_xy_variant {
319 ($name:ident, $ctype:ident, $(#[$($attrss:tt)*])*) => {
320 $(#[$($attrss)*])*
321 pub fn $name(x: Option<Expr>, y: Option<Expr>) -> Arc<CalculationRule> {
322 std::sync::Arc::new(CalculationRule { calculation_type: CalculationType::$ctype{ x, y } } )
323 }
324 }
325}
326
327calc_empty_variant!(dfq_count, Count, #[doc = r#"Creates a rule that counts the number of rows in a column.
328
329# Examples
330
331```
332use datafusion_quality::rules::table::dfq_count;
333use datafusion_quality::RuleSet;
334
335// Create a rule to count the number of rows in the age column
336let rule = dfq_count();
337let mut ruleset = RuleSet::new();
338ruleset.with_table_rule("age", rule, None);
339```"#]);
340calc_empty_variant!(dfq_count_distinct, CountDistinct, #[doc = r#"Creates a rule that counts the number of distinct values in a column.
341
342# Examples
343
344```
345use datafusion_quality::rules::table::dfq_count_distinct;
346use datafusion_quality::RuleSet;
347
348// Create a rule to count distinct values in the age column
349let rule = dfq_count_distinct();
350let mut ruleset = RuleSet::new();
351ruleset.with_table_rule("age", rule, None);
352```"#]);
353calc_empty_variant!(dfq_avg, Avg, #[doc = r#"Creates a rule that calculates the average value of a column.
354
355# Examples
356
357```
358use datafusion_quality::rules::table::dfq_avg;
359use datafusion_quality::RuleSet;
360
361// Create a rule to calculate the average of the age column
362let rule = dfq_avg();
363let mut ruleset = RuleSet::new();
364ruleset.with_table_rule("age", rule, None);
365```"#]);
366calc_empty_variant!(dfq_stddev, StdDev, #[doc = r#"Creates a rule that calculates the standard deviation of a column.
367
368# Examples
369
370```
371use datafusion_quality::rules::table::dfq_stddev;
372use datafusion_quality::RuleSet;
373
374// Create a rule to calculate the standard deviation of the age column
375let rule = dfq_stddev();
376let mut ruleset = RuleSet::new();
377ruleset.with_table_rule("age", rule, None);
378```"#]);
379calc_empty_variant!(dfq_max, Max, #[doc = r#"Creates a rule that calculates the maximum value of a column.
380
381# Examples
382
383```
384use datafusion_quality::rules::table::dfq_max;
385use datafusion_quality::RuleSet;
386
387// Create a rule to find the maximum value in the age column
388let rule = dfq_max();
389let mut ruleset = RuleSet::new();
390ruleset.with_table_rule("age", rule, None);
391```"#]);
392calc_empty_variant!(dfq_min, Min, #[doc = r#"Creates a rule that calculates the minimum value of a column.
393
394# Examples
395
396```
397use datafusion_quality::rules::table::dfq_min;
398use datafusion_quality::RuleSet;
399
400// Create a rule to find the minimum value in the age column
401let rule = dfq_min();
402let mut ruleset = RuleSet::new();
403ruleset.with_table_rule("age", rule, None);
404```"#]);
405calc_empty_variant!(dfq_sum, Sum, #[doc = r#"Creates a rule that calculates the sum of a column.
406
407# Examples
408
409```
410use datafusion_quality::rules::table::dfq_sum;
411use datafusion_quality::RuleSet;
412
413// Create a rule to calculate the sum of the age column
414let rule = dfq_sum();
415let mut ruleset = RuleSet::new();
416ruleset.with_table_rule("age", rule, None);
417```"#]);
418calc_empty_variant!(dfq_median, Median, #[doc = r#"Creates a rule that calculates the median of a column.
419
420# Examples
421
422```
423use datafusion_quality::rules::table::dfq_median;
424use datafusion_quality::RuleSet;
425
426// Create a rule to calculate the median of the age column
427let rule = dfq_median();
428let mut ruleset = RuleSet::new();
429ruleset.with_table_rule("age", rule, None);
430```"#]);
431calc_empty_variant!(dfq_last_value, LastValue, #[doc = r#"Creates a rule that calculates the last value of a column.
432
433# Examples
434
435```
436use datafusion_quality::rules::table::dfq_last_value;
437use datafusion_quality::RuleSet;
438
439// Create a rule to get the last value in the age column
440let rule = dfq_last_value();
441let mut ruleset = RuleSet::new();
442ruleset.with_table_rule("age", rule, None);
443```"#]);
444calc_empty_variant!(dfq_stddev_pop, StddevPop, #[doc = r#"Creates a rule that calculates the population standard deviation of a column.
445
446# Examples
447
448```
449use datafusion_quality::rules::table::dfq_stddev_pop;
450use datafusion_quality::RuleSet;
451
452// Create a rule to calculate the population standard deviation of the age column
453let rule = dfq_stddev_pop();
454let mut ruleset = RuleSet::new();
455ruleset.with_table_rule("age", rule, None);
456```"#]);
457calc_empty_variant!(dfq_var_pop, VarPop, #[doc = r#"Creates a rule that calculates the population variance of a column.
458
459# Examples
460
461```
462use datafusion_quality::rules::table::dfq_var_pop;
463use datafusion_quality::RuleSet;
464
465// Create a rule to calculate the population variance of the age column
466let rule = dfq_var_pop();
467let mut ruleset = RuleSet::new();
468ruleset.with_table_rule("age", rule, None);
469```"#]);
470calc_empty_variant!(dfq_var_samp, VarSamp, #[doc = r#"Creates a rule that calculates the sample variance of a column.
471
472# Examples
473
474```
475use datafusion_quality::rules::table::dfq_var_samp;
476use datafusion_quality::RuleSet;
477
478// Create a rule to calculate the sample variance of the age column
479let rule = dfq_var_samp();
480let mut ruleset = RuleSet::new();
481ruleset.with_table_rule("age", rule, None);
482```"#]);
483calc_xy_variant!(dfq_covar_pop, CovarPop, #[doc = r#"Creates a rule that calculates the population covariance of two columns.
484
485# Examples
486
487```
488use datafusion_quality::rules::table::dfq_covar_pop;
489use datafusion_quality::RuleSet;
490use datafusion::prelude::*;
491
492// Create a rule to calculate the population covariance between age and score columns
493let rule = dfq_covar_pop(Some(col("age")), Some(col("score")));
494let mut ruleset = RuleSet::new();
495ruleset.with_table_rule("age", rule, None);
496```"#]);
497calc_xy_variant!(dfq_covar_samp, CovarSamp, #[doc = r#"Creates a rule that calculates the sample covariance of two columns.
498
499# Examples
500
501```
502use datafusion_quality::rules::table::dfq_covar_samp;
503use datafusion_quality::RuleSet;
504use datafusion::prelude::*;
505
506// Create a rule to calculate the sample covariance between age and score columns
507let rule = dfq_covar_samp(Some(col("age")), Some(col("score")));
508let mut ruleset = RuleSet::new();
509ruleset.with_table_rule("age", rule, None);
510```"#]);
511calc_xy_variant!(dfq_regr_avgx, RegrAvgX, #[doc = r#"Creates a rule that calculates the average of x values in a column.
512
513# Examples
514
515```
516use datafusion_quality::rules::table::dfq_regr_avgx;
517use datafusion_quality::RuleSet;
518use datafusion::prelude::*;
519
520// Create a rule to calculate the average of x values between age and score columns
521let rule = dfq_regr_avgx(Some(col("age")), Some(col("score")));
522let mut ruleset = RuleSet::new();
523ruleset.with_table_rule("age", rule, None);
524```"#]);
525calc_xy_variant!(dfq_regr_avgy, RegrAvgY, #[doc = r#"Creates a rule that calculates the average of y values in a column.
526
527# Examples
528
529```
530use datafusion_quality::rules::table::dfq_regr_avgy;
531use datafusion_quality::RuleSet;
532use datafusion::prelude::*;
533
534// Create a rule to calculate the average of y values between age and score columns
535let rule = dfq_regr_avgy(Some(col("age")), Some(col("score")));
536let mut ruleset = RuleSet::new();
537ruleset.with_table_rule("age", rule, None);
538```"#]);
539calc_xy_variant!(dfq_regr_count, RegrCount, #[doc = r#"Creates a rule that calculates the number of rows in a column.
540
541# Examples
542
543```
544use datafusion_quality::rules::table::dfq_regr_count;
545use datafusion_quality::RuleSet;
546use datafusion::prelude::*;
547
548// Create a rule to count rows between age and score columns
549let rule = dfq_regr_count(Some(col("age")), Some(col("score")));
550let mut ruleset = RuleSet::new();
551ruleset.with_table_rule("age", rule, None);
552```"#]);
553calc_xy_variant!(dfq_regr_intercept, RegrIntercept, #[doc = r#"Creates a rule that calculates the intercept of a linear regression.
554
555# Examples
556
557```
558use datafusion_quality::rules::table::dfq_regr_intercept;
559use datafusion_quality::RuleSet;
560use datafusion::prelude::*;
561
562// Create a rule to calculate the intercept between age and score columns
563let rule = dfq_regr_intercept(Some(col("age")), Some(col("score")));
564let mut ruleset = RuleSet::new();
565ruleset.with_table_rule("age", rule, None);
566```"#]);
567calc_xy_variant!(dfq_regr_r2, RegrR2, #[doc = r#"Creates a rule that calculates the R-squared value of a linear regression.
568
569# Examples
570
571```
572use datafusion_quality::rules::table::dfq_regr_r2;
573use datafusion_quality::RuleSet;
574use datafusion::prelude::*;
575
576// Create a rule to calculate the R-squared value between age and score columns
577let rule = dfq_regr_r2(Some(col("age")), Some(col("score")));
578let mut ruleset = RuleSet::new();
579ruleset.with_table_rule("age", rule, None);
580```"#]);
581calc_xy_variant!(dfq_regr_slope, RegrSlope, #[doc = r#"Creates a rule that calculates the slope of a linear regression.
582
583# Examples
584
585```
586use datafusion_quality::rules::table::dfq_regr_slope;
587use datafusion_quality::RuleSet;
588use datafusion::prelude::*;
589
590// Create a rule to calculate the slope between age and score columns
591let rule = dfq_regr_slope(Some(col("age")), Some(col("score")));
592let mut ruleset = RuleSet::new();
593ruleset.with_table_rule("age", rule, None);
594```"#]);
595calc_xy_variant!(dfq_regr_sxx, RegrSxx, #[doc = r#"Creates a rule that calculates the sum of squared deviations from the mean for x values.
596
597# Examples
598
599```
600use datafusion_quality::rules::table::dfq_regr_sxx;
601use datafusion_quality::RuleSet;
602use datafusion::prelude::*;
603
604// Create a rule to calculate the sum of squared deviations for x values between age and score columns
605let rule = dfq_regr_sxx(Some(col("age")), Some(col("score")));
606let mut ruleset = RuleSet::new();
607ruleset.with_table_rule("age", rule, None);
608```"#]);
609calc_xy_variant!(dfq_regr_sxy, RegrSxy, #[doc = r#"Creates a rule that calculates the sum of the products of deviations from the mean for x and y values.
610
611# Examples
612
613```
614use datafusion_quality::rules::table::dfq_regr_sxy;
615use datafusion_quality::RuleSet;
616use datafusion::prelude::*;
617
618// Create a rule to calculate the sum of products of deviations between age and score columns
619let rule = dfq_regr_sxy(Some(col("age")), Some(col("score")));
620let mut ruleset = RuleSet::new();
621ruleset.with_table_rule("age", rule, None);
622```"#]);
623calc_xy_variant!(dfq_regr_syy, RegrSyy, #[doc = r#"Creates a rule that calculates the sum of squared deviations from the mean for y values.
624
625# Examples
626
627```
628use datafusion_quality::rules::table::dfq_regr_syy;
629use datafusion_quality::RuleSet;
630use datafusion::prelude::*;
631
632// Create a rule to calculate the sum of squared deviations for y values between age and score columns
633let rule = dfq_regr_syy(Some(col("age")), Some(col("score")));
634let mut ruleset = RuleSet::new();
635ruleset.with_table_rule("age", rule, None);
636```"#]);
637
638pub fn dfq_nth_value(n: i64, sort_exprs: Option<Vec<SortExpr>>) -> Arc<CalculationRule> {
657 std::sync::Arc::new(CalculationRule {
658 calculation_type: CalculationType::NthValue(n, sort_exprs),
659 })
660}
661
662pub fn dfq_first_value(sort_exprs: Option<Vec<SortExpr>>) -> Arc<CalculationRule> {
680 std::sync::Arc::new(CalculationRule {
681 calculation_type: CalculationType::FirstValue(sort_exprs),
682 })
683}
684
685#[derive(Debug, Clone, Default)]
686pub struct CustomAggregationRuleBuilder {
687 aggregation: Expr,
688 rule_name: String,
689 group_by_exprs: Option<Vec<Expr>>,
690 aggregate_exprs: Option<Vec<Expr>>,
691 order_by: Option<Vec<SortExpr>>,
692 window_exprs: Option<Vec<Expr>>,
693 filter: Option<Expr>,
694}
695
696impl CustomAggregationRuleBuilder {
697 pub fn new(aggregation: Expr, rule_name: String) -> Self {
698 Self {
699 aggregation,
700 rule_name,
701 ..Default::default()
702 }
703 }
704
705 pub fn with_group_by(mut self, group_by_exprs: Vec<Expr>) -> Self {
706 self.group_by_exprs = Some(group_by_exprs);
707 self
708 }
709
710 pub fn with_aggregate_exprs(mut self, aggregate_exprs: Vec<Expr>) -> Self {
711 self.aggregate_exprs = Some(aggregate_exprs);
712 self
713 }
714
715 pub fn with_order_by(mut self, order_by: Vec<SortExpr>) -> Self {
716 self.order_by = Some(order_by);
717 self
718 }
719
720 pub fn with_window_exprs(mut self, window_exprs: Vec<Expr>) -> Self {
721 self.window_exprs = Some(window_exprs);
722 self
723 }
724
725 pub fn with_filter(mut self, filter: Expr) -> Self {
726 self.filter = Some(filter);
727 self
728 }
729
730 pub fn build(self) -> Arc<CustomAggregationRule> {
731 std::sync::Arc::new(CustomAggregationRule {
732 aggregation: self.aggregation,
733 rule_name: self.rule_name,
734 group_by_exprs: self.group_by_exprs,
735 aggregate_exprs: self.aggregate_exprs,
736 order_by: self.order_by,
737 window_exprs: self.window_exprs,
738 filter: self.filter,
739 })
740 }
741}
742
743#[derive(Debug, Clone, Default)]
745pub struct CustomAggregationRule {
746 aggregation: Expr,
747 group_by_exprs: Option<Vec<Expr>>,
748 aggregate_exprs: Option<Vec<Expr>>,
749 order_by: Option<Vec<SortExpr>>,
750 window_exprs: Option<Vec<Expr>>,
751 filter: Option<Expr>,
752 rule_name: String,
753}
754
755impl CustomAggregationRule {
756 pub fn builder(aggregation: Expr, rule_name: String) -> CustomAggregationRuleBuilder {
757 CustomAggregationRuleBuilder::new(aggregation, rule_name)
758 }
759}
760
761impl TableRule for CustomAggregationRule {
762 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
763 let mut subquery = df.clone();
764 if let Some(filter) = self.filter.clone() {
765 subquery = subquery.filter(filter)?;
766 }
767
768 match (self.group_by_exprs.clone(), self.aggregate_exprs.clone()) {
769 (Some(group_by), Some(aggregate)) => {
770 subquery = subquery.aggregate(group_by, aggregate)?;
771 }
772 (None, Some(aggregate)) => {
773 subquery = subquery.aggregate(vec![], aggregate)?;
774 }
775 _ => {
776 return Err(ValidationError::Configuration {
777 message: "Group by requires aggregate expressions".to_string(),
778 });
779 }
780 }
781
782 if let Some(window_exprs) = self.window_exprs.clone() {
783 subquery = subquery.window(window_exprs)?;
784 }
785 if let Some(order_by) = self.order_by.clone() {
786 subquery = subquery.sort(order_by)?;
787 }
788 subquery = subquery.select(vec![self.aggregation.clone()])?;
789
790 let subq_expr = Expr::ScalarSubquery(Subquery {
791 subquery: Arc::new(subquery.logical_plan().clone()),
792 outer_ref_columns: vec![],
793 });
794 df.with_column(&self.new_column_name(column_name), subq_expr)
795 .context(DataFusionSnafu)
796 }
797
798 fn name(&self) -> &str {
799 &self.rule_name
800 }
801
802 fn new_column_name(&self, column_name: &str) -> String {
803 format!("{}_{}", column_name, self.rule_name)
804 }
805
806 fn description(&self) -> &str {
807 "Applies a custom aggregation across the entire table"
808 }
809}
810
811pub fn dfq_custom_agg(aggregation: Expr, rule_name: String) -> Arc<CustomAggregationRule> {
812 CustomAggregationRule::builder(aggregation, rule_name).build()
813}
814
815#[cfg(test)]
816mod tests {
817 use super::*;
818 use arrow::record_batch::RecordBatch;
819 use datafusion::arrow::array::{Float64Array, Int32Array, StringArray};
820 use datafusion::arrow::datatypes::{DataType, Field, Schema};
821 use datafusion::assert_batches_eq;
822 use std::sync::Arc;
823
824 async fn create_test_df() -> DataFrame {
825 let schema = Schema::new(vec![
826 Field::new("id", DataType::Int32, false),
827 Field::new("name", DataType::Utf8, true),
828 Field::new("age", DataType::Int32, true),
829 Field::new("score", DataType::Float64, true),
830 ]);
831
832 let id_data = Int32Array::from(vec![1, 2, 3, 4, 5]);
833 let name_data = StringArray::from(vec![
834 Some("Alice"),
835 Some("Bob"),
836 None,
837 Some("Charlie"),
838 Some("Dave"),
839 ]);
840 let age_data = Int32Array::from(vec![Some(25), Some(30), Some(15), Some(40), Some(25)]);
841 let score_data = Float64Array::from(vec![
842 Some(85.5),
843 Some(92.0),
844 Some(78.5),
845 Some(95.0),
846 Some(88.5),
847 ]);
848
849 let batch = RecordBatch::try_new(
850 Arc::new(schema),
851 vec![
852 Arc::new(id_data),
853 Arc::new(name_data),
854 Arc::new(age_data),
855 Arc::new(score_data),
856 ],
857 )
858 .unwrap();
859
860 let ctx = SessionContext::new();
861 ctx.read_batch(batch).unwrap()
862 }
863
864 #[tokio::test]
865 async fn test_null_count_rule() {
866 let df = create_test_df().await;
867 let rule = NullCountRule::default();
868 let result = rule.apply(df, "name").unwrap();
869
870 let expected = vec![
871 "+----+---------+-----+-------+-----------------+",
872 "| id | name | age | score | name_null_count |",
873 "+----+---------+-----+-------+-----------------+",
874 "| 1 | Alice | 25 | 85.5 | 1 |",
875 "| 2 | Bob | 30 | 92.0 | 1 |",
876 "| 3 | | 15 | 78.5 | 1 |",
877 "| 4 | Charlie | 40 | 95.0 | 1 |",
878 "| 5 | Dave | 25 | 88.5 | 1 |",
879 "+----+---------+-----+-------+-----------------+",
880 ];
881
882 assert_batches_eq!(&expected, &result.collect().await.unwrap());
883
884 let df = create_test_df().await;
886 let rule = NullCountRule {
887 negated: Some(true),
888 };
889 let result = rule.apply(df, "name").unwrap();
890
891 let expected = vec![
892 "+----+---------+-----+-------+---------------------+",
893 "| id | name | age | score | name_not_null_count |",
894 "+----+---------+-----+-------+---------------------+",
895 "| 1 | Alice | 25 | 85.5 | 4 |",
896 "| 2 | Bob | 30 | 92.0 | 4 |",
897 "| 3 | | 15 | 78.5 | 4 |",
898 "| 4 | Charlie | 40 | 95.0 | 4 |",
899 "| 5 | Dave | 25 | 88.5 | 4 |",
900 "+----+---------+-----+-------+---------------------+",
901 ];
902
903 assert_batches_eq!(&expected, &result.collect().await.unwrap());
904 }
905
906 #[tokio::test]
907 async fn test_count_rule() {
908 let df = create_test_df().await;
909 let rule = dfq_count();
910 let result = rule.apply(df, "age").unwrap();
911
912 let expected = vec![
913 "+----+---------+-----+-------+-----------+",
914 "| id | name | age | score | age_count |",
915 "+----+---------+-----+-------+-----------+",
916 "| 1 | Alice | 25 | 85.5 | 5 |",
917 "| 2 | Bob | 30 | 92.0 | 5 |",
918 "| 3 | | 15 | 78.5 | 5 |",
919 "| 4 | Charlie | 40 | 95.0 | 5 |",
920 "| 5 | Dave | 25 | 88.5 | 5 |",
921 "+----+---------+-----+-------+-----------+",
922 ];
923
924 assert_batches_eq!(&expected, &result.collect().await.unwrap());
925 }
926
927 #[tokio::test]
928 async fn test_count_distinct_rule() {
929 let df = create_test_df().await;
930 let rule = dfq_count_distinct();
931 let result = rule.apply(df, "age").unwrap();
932
933 let expected = vec![
934 "+----+---------+-----+-------+-------------------+",
935 "| id | name | age | score | age_countdistinct |",
936 "+----+---------+-----+-------+-------------------+",
937 "| 1 | Alice | 25 | 85.5 | 4 |",
938 "| 2 | Bob | 30 | 92.0 | 4 |",
939 "| 3 | | 15 | 78.5 | 4 |",
940 "| 4 | Charlie | 40 | 95.0 | 4 |",
941 "| 5 | Dave | 25 | 88.5 | 4 |",
942 "+----+---------+-----+-------+-------------------+",
943 ];
944
945 assert_batches_eq!(&expected, &result.collect().await.unwrap());
946 }
947
948 #[tokio::test]
949 async fn test_avg_rule() {
950 let df = create_test_df().await;
951 let rule = dfq_avg();
952 let result = rule.apply(df, "score").unwrap();
953
954 let expected = vec![
955 "+----+---------+-----+-------+-----------+",
956 "| id | name | age | score | score_avg |",
957 "+----+---------+-----+-------+-----------+",
958 "| 1 | Alice | 25 | 85.5 | 87.9 |",
959 "| 2 | Bob | 30 | 92.0 | 87.9 |",
960 "| 3 | | 15 | 78.5 | 87.9 |",
961 "| 4 | Charlie | 40 | 95.0 | 87.9 |",
962 "| 5 | Dave | 25 | 88.5 | 87.9 |",
963 "+----+---------+-----+-------+-----------+",
964 ];
965
966 assert_batches_eq!(&expected, &result.collect().await.unwrap());
967 }
968
969 #[tokio::test]
970 async fn test_stddev_rule() {
971 let df = create_test_df().await;
972 let rule = dfq_stddev();
973 let result = rule.apply(df, "score").unwrap();
974
975 let expected = vec![
976 "+----+---------+-----+-------+-------------------+",
977 "| id | name | age | score | score_stddev |",
978 "+----+---------+-----+-------+-------------------+",
979 "| 1 | Alice | 25 | 85.5 | 6.358065743604732 |",
980 "| 2 | Bob | 30 | 92.0 | 6.358065743604732 |",
981 "| 3 | | 15 | 78.5 | 6.358065743604732 |",
982 "| 4 | Charlie | 40 | 95.0 | 6.358065743604732 |",
983 "| 5 | Dave | 25 | 88.5 | 6.358065743604732 |",
984 "+----+---------+-----+-------+-------------------+",
985 ];
986
987 assert_batches_eq!(&expected, &result.collect().await.unwrap());
988 }
989
990 #[tokio::test]
991 async fn test_max_rule() {
992 let df = create_test_df().await;
993 let rule = dfq_max();
994 let result = rule.apply(df, "score").unwrap();
995
996 let expected = vec![
997 "+----+---------+-----+-------+-----------+",
998 "| id | name | age | score | score_max |",
999 "+----+---------+-----+-------+-----------+",
1000 "| 1 | Alice | 25 | 85.5 | 95.0 |",
1001 "| 2 | Bob | 30 | 92.0 | 95.0 |",
1002 "| 3 | | 15 | 78.5 | 95.0 |",
1003 "| 4 | Charlie | 40 | 95.0 | 95.0 |",
1004 "| 5 | Dave | 25 | 88.5 | 95.0 |",
1005 "+----+---------+-----+-------+-----------+",
1006 ];
1007
1008 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1009 }
1010
1011 #[tokio::test]
1012 async fn test_min_rule() {
1013 let df = create_test_df().await;
1014 let rule = dfq_min();
1015 let result = rule.apply(df, "score").unwrap();
1016
1017 let expected = vec![
1018 "+----+---------+-----+-------+-----------+",
1019 "| id | name | age | score | score_min |",
1020 "+----+---------+-----+-------+-----------+",
1021 "| 1 | Alice | 25 | 85.5 | 78.5 |",
1022 "| 2 | Bob | 30 | 92.0 | 78.5 |",
1023 "| 3 | | 15 | 78.5 | 78.5 |",
1024 "| 4 | Charlie | 40 | 95.0 | 78.5 |",
1025 "| 5 | Dave | 25 | 88.5 | 78.5 |",
1026 "+----+---------+-----+-------+-----------+",
1027 ];
1028
1029 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1030 }
1031
1032 #[tokio::test]
1033 async fn test_sum_rule() {
1034 let df = create_test_df().await;
1035 let rule = dfq_sum();
1036 let result = rule.apply(df, "score").unwrap();
1037
1038 let expected = vec![
1039 "+----+---------+-----+-------+-----------+",
1040 "| id | name | age | score | score_sum |",
1041 "+----+---------+-----+-------+-----------+",
1042 "| 1 | Alice | 25 | 85.5 | 439.5 |",
1043 "| 2 | Bob | 30 | 92.0 | 439.5 |",
1044 "| 3 | | 15 | 78.5 | 439.5 |",
1045 "| 4 | Charlie | 40 | 95.0 | 439.5 |",
1046 "| 5 | Dave | 25 | 88.5 | 439.5 |",
1047 "+----+---------+-----+-------+-----------+",
1048 ];
1049
1050 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1051 }
1052
1053 #[tokio::test]
1054 async fn test_median_rule() {
1055 let df = create_test_df().await;
1056 let rule = dfq_median();
1057 let result = rule.apply(df, "score").unwrap();
1058
1059 let expected = vec![
1060 "+----+---------+-----+-------+--------------+",
1061 "| id | name | age | score | score_median |",
1062 "+----+---------+-----+-------+--------------+",
1063 "| 1 | Alice | 25 | 85.5 | 88.5 |",
1064 "| 2 | Bob | 30 | 92.0 | 88.5 |",
1065 "| 3 | | 15 | 78.5 | 88.5 |",
1066 "| 4 | Charlie | 40 | 95.0 | 88.5 |",
1067 "| 5 | Dave | 25 | 88.5 | 88.5 |",
1068 "+----+---------+-----+-------+--------------+",
1069 ];
1070
1071 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1072 }
1073
1074 #[tokio::test]
1075 async fn test_last_value_rule() {
1076 let df = create_test_df().await;
1077 let rule = dfq_last_value();
1078 let result = rule.apply(df, "score").unwrap();
1079
1080 let expected = vec![
1081 "+----+---------+-----+-------+-----------------+",
1082 "| id | name | age | score | score_lastvalue |",
1083 "+----+---------+-----+-------+-----------------+",
1084 "| 1 | Alice | 25 | 85.5 | 88.5 |",
1085 "| 2 | Bob | 30 | 92.0 | 88.5 |",
1086 "| 3 | | 15 | 78.5 | 88.5 |",
1087 "| 4 | Charlie | 40 | 95.0 | 88.5 |",
1088 "| 5 | Dave | 25 | 88.5 | 88.5 |",
1089 "+----+---------+-----+-------+-----------------+",
1090 ];
1091
1092 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1093 }
1094
1095 #[tokio::test]
1096 async fn test_stddev_pop_rule() {
1097 let df = create_test_df().await;
1098 let rule = dfq_stddev_pop();
1099 let result = rule.apply(df, "score").unwrap();
1100
1101 let expected = vec![
1102 "+----+---------+-----+-------+-------------------+",
1103 "| id | name | age | score | score_stddevpop |",
1104 "+----+---------+-----+-------+-------------------+",
1105 "| 1 | Alice | 25 | 85.5 | 5.686826883245172 |",
1106 "| 2 | Bob | 30 | 92.0 | 5.686826883245172 |",
1107 "| 3 | | 15 | 78.5 | 5.686826883245172 |",
1108 "| 4 | Charlie | 40 | 95.0 | 5.686826883245172 |",
1109 "| 5 | Dave | 25 | 88.5 | 5.686826883245172 |",
1110 "+----+---------+-----+-------+-------------------+",
1111 ];
1112
1113 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1114 }
1115
1116 #[tokio::test]
1117 async fn test_var_pop_rule() {
1118 let df = create_test_df().await;
1119 let rule = dfq_var_pop();
1120 let result = rule.apply(df, "score").unwrap();
1121
1122 let expected = vec![
1123 "+----+---------+-----+-------+--------------------+",
1124 "| id | name | age | score | score_varpop |",
1125 "+----+---------+-----+-------+--------------------+",
1126 "| 1 | Alice | 25 | 85.5 | 32.339999999999996 |",
1127 "| 2 | Bob | 30 | 92.0 | 32.339999999999996 |",
1128 "| 3 | | 15 | 78.5 | 32.339999999999996 |",
1129 "| 4 | Charlie | 40 | 95.0 | 32.339999999999996 |",
1130 "| 5 | Dave | 25 | 88.5 | 32.339999999999996 |",
1131 "+----+---------+-----+-------+--------------------+",
1132 ];
1133
1134 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1135 }
1136
1137 #[tokio::test]
1138 async fn test_var_samp_rule() {
1139 let df = create_test_df().await;
1140 let rule = dfq_var_samp();
1141 let result = rule.apply(df, "score").unwrap();
1142
1143 let expected = vec![
1144 "+----+---------+-----+-------+---------------+",
1145 "| id | name | age | score | score_varsamp |",
1146 "+----+---------+-----+-------+---------------+",
1147 "| 1 | Alice | 25 | 85.5 | 40.425 |",
1148 "| 2 | Bob | 30 | 92.0 | 40.425 |",
1149 "| 3 | | 15 | 78.5 | 40.425 |",
1150 "| 4 | Charlie | 40 | 95.0 | 40.425 |",
1151 "| 5 | Dave | 25 | 88.5 | 40.425 |",
1152 "+----+---------+-----+-------+---------------+",
1153 ];
1154
1155 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1156 }
1157
1158 #[tokio::test]
1159 async fn test_covar_pop_rule() {
1160 let df = create_test_df().await;
1161 let rule = dfq_covar_pop(Some(col("age")), None);
1162 let result = rule.apply(df, "score").unwrap();
1163
1164 let expected = vec![
1165 "+----+---------+-----+-------+-------------------+",
1166 "| id | name | age | score | score_covarpop |",
1167 "+----+---------+-----+-------+-------------------+",
1168 "| 1 | Alice | 25 | 85.5 | 44.20000000000001 |",
1169 "| 2 | Bob | 30 | 92.0 | 44.20000000000001 |",
1170 "| 3 | | 15 | 78.5 | 44.20000000000001 |",
1171 "| 4 | Charlie | 40 | 95.0 | 44.20000000000001 |",
1172 "| 5 | Dave | 25 | 88.5 | 44.20000000000001 |",
1173 "+----+---------+-----+-------+-------------------+",
1174 ];
1175
1176 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1177 }
1178
1179 #[tokio::test]
1180 async fn test_covar_samp_rule() {
1181 let df = create_test_df().await;
1182 let rule = dfq_covar_samp(Some(col("age")), None);
1183 let result = rule.apply(df.clone(), "score").unwrap();
1184
1185 let expected = vec![
1186 "+----+---------+-----+-------+--------------------+",
1187 "| id | name | age | score | score_covarsamp |",
1188 "+----+---------+-----+-------+--------------------+",
1189 "| 1 | Alice | 25 | 85.5 | 55.250000000000014 |",
1190 "| 2 | Bob | 30 | 92.0 | 55.250000000000014 |",
1191 "| 3 | | 15 | 78.5 | 55.250000000000014 |",
1192 "| 4 | Charlie | 40 | 95.0 | 55.250000000000014 |",
1193 "| 5 | Dave | 25 | 88.5 | 55.250000000000014 |",
1194 "+----+---------+-----+-------+--------------------+",
1195 ];
1196
1197 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1198
1199 let rule = dfq_covar_samp(None, Some(col("age")));
1200 let result = rule.apply(df, "score").unwrap();
1201
1202 let expected = vec![
1203 "+----+---------+-----+-------+--------------------+",
1204 "| id | name | age | score | score_covarsamp |",
1205 "+----+---------+-----+-------+--------------------+",
1206 "| 1 | Alice | 25 | 85.5 | 55.249999999999986 |",
1207 "| 2 | Bob | 30 | 92.0 | 55.249999999999986 |",
1208 "| 3 | | 15 | 78.5 | 55.249999999999986 |",
1209 "| 4 | Charlie | 40 | 95.0 | 55.249999999999986 |",
1210 "| 5 | Dave | 25 | 88.5 | 55.249999999999986 |",
1211 "+----+---------+-----+-------+--------------------+",
1212 ];
1213
1214 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1215 }
1216
1217 #[tokio::test]
1218 async fn test_regr_avgx_rule() {
1219 let df = create_test_df().await;
1220 let rule = dfq_regr_avgx(Some(col("age")), None);
1221 let result = rule.apply(df.clone(), "score").unwrap();
1222
1223 let expected = vec![
1224 "+----+---------+-----+-------+----------------+",
1225 "| id | name | age | score | score_regravgx |",
1226 "+----+---------+-----+-------+----------------+",
1227 "| 1 | Alice | 25 | 85.5 | 27.0 |",
1228 "| 2 | Bob | 30 | 92.0 | 27.0 |",
1229 "| 3 | | 15 | 78.5 | 27.0 |",
1230 "| 4 | Charlie | 40 | 95.0 | 27.0 |",
1231 "| 5 | Dave | 25 | 88.5 | 27.0 |",
1232 "+----+---------+-----+-------+----------------+",
1233 ];
1234
1235 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1236
1237 let rule = dfq_regr_avgx(None, Some(col("age")));
1238 let result = rule.apply(df, "score").unwrap();
1239
1240 let expected = vec![
1241 "+----+---------+-----+-------+----------------+",
1242 "| id | name | age | score | score_regravgx |",
1243 "+----+---------+-----+-------+----------------+",
1244 "| 1 | Alice | 25 | 85.5 | 87.9 |",
1245 "| 2 | Bob | 30 | 92.0 | 87.9 |",
1246 "| 3 | | 15 | 78.5 | 87.9 |",
1247 "| 4 | Charlie | 40 | 95.0 | 87.9 |",
1248 "| 5 | Dave | 25 | 88.5 | 87.9 |",
1249 "+----+---------+-----+-------+----------------+",
1250 ];
1251
1252 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1253 }
1254
1255 #[tokio::test]
1256 async fn test_regr_avgy_rule() {
1257 let df = create_test_df().await;
1258 let rule = dfq_regr_avgy(Some(col("age")), None);
1259 let result = rule.apply(df.clone(), "score").unwrap();
1260
1261 let expected = vec![
1262 "+----+---------+-----+-------+----------------+",
1263 "| id | name | age | score | score_regravgy |",
1264 "+----+---------+-----+-------+----------------+",
1265 "| 1 | Alice | 25 | 85.5 | 87.9 |",
1266 "| 2 | Bob | 30 | 92.0 | 87.9 |",
1267 "| 3 | | 15 | 78.5 | 87.9 |",
1268 "| 4 | Charlie | 40 | 95.0 | 87.9 |",
1269 "| 5 | Dave | 25 | 88.5 | 87.9 |",
1270 "+----+---------+-----+-------+----------------+",
1271 ];
1272
1273 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1274
1275 let rule = dfq_regr_avgy(None, Some(col("age")));
1276 let result = rule.apply(df, "score").unwrap();
1277
1278 let expected = vec![
1279 "+----+---------+-----+-------+----------------+",
1280 "| id | name | age | score | score_regravgy |",
1281 "+----+---------+-----+-------+----------------+",
1282 "| 1 | Alice | 25 | 85.5 | 27.0 |",
1283 "| 2 | Bob | 30 | 92.0 | 27.0 |",
1284 "| 3 | | 15 | 78.5 | 27.0 |",
1285 "| 4 | Charlie | 40 | 95.0 | 27.0 |",
1286 "| 5 | Dave | 25 | 88.5 | 27.0 |",
1287 "+----+---------+-----+-------+----------------+",
1288 ];
1289
1290 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1291 }
1292
1293 #[tokio::test]
1294 async fn test_regr_count_rule() {
1295 let df = create_test_df().await;
1296 let rule = dfq_regr_count(Some(col("age")), None);
1297 let result = rule.apply(df.clone(), "score").unwrap();
1298
1299 let expected = vec![
1300 "+----+---------+-----+-------+-----------------+",
1301 "| id | name | age | score | score_regrcount |",
1302 "+----+---------+-----+-------+-----------------+",
1303 "| 1 | Alice | 25 | 85.5 | 5 |",
1304 "| 2 | Bob | 30 | 92.0 | 5 |",
1305 "| 3 | | 15 | 78.5 | 5 |",
1306 "| 4 | Charlie | 40 | 95.0 | 5 |",
1307 "| 5 | Dave | 25 | 88.5 | 5 |",
1308 "+----+---------+-----+-------+-----------------+",
1309 ];
1310
1311 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1312
1313 let rule = dfq_regr_count(None, Some(col("age")));
1314 let result = rule.apply(df, "score").unwrap();
1315
1316 let expected = vec![
1317 "+----+---------+-----+-------+-----------------+",
1318 "| id | name | age | score | score_regrcount |",
1319 "+----+---------+-----+-------+-----------------+",
1320 "| 1 | Alice | 25 | 85.5 | 5 |",
1321 "| 2 | Bob | 30 | 92.0 | 5 |",
1322 "| 3 | | 15 | 78.5 | 5 |",
1323 "| 4 | Charlie | 40 | 95.0 | 5 |",
1324 "| 5 | Dave | 25 | 88.5 | 5 |",
1325 "+----+---------+-----+-------+-----------------+",
1326 ];
1327
1328 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1329 }
1330
1331 #[tokio::test]
1332 async fn test_regr_intercept_rule() {
1333 let df = create_test_df().await;
1334 let rule = dfq_regr_intercept(Some(col("age")), None);
1335 let result = rule.apply(df.clone(), "score").unwrap();
1336
1337 let expected = vec![
1338 "+----+---------+-----+-------+---------------------+",
1339 "| id | name | age | score | score_regrintercept |",
1340 "+----+---------+-----+-------+---------------------+",
1341 "| 1 | Alice | 25 | 85.5 | 69.81818181818183 |",
1342 "| 2 | Bob | 30 | 92.0 | 69.81818181818183 |",
1343 "| 3 | | 15 | 78.5 | 69.81818181818183 |",
1344 "| 4 | Charlie | 40 | 95.0 | 69.81818181818183 |",
1345 "| 5 | Dave | 25 | 88.5 | 69.81818181818183 |",
1346 "+----+---------+-----+-------+---------------------+",
1347 ];
1348
1349 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1350
1351 let rule = dfq_regr_intercept(None, Some(col("age")));
1352 let result = rule.apply(df, "score").unwrap();
1353
1354 let expected = vec![
1355 "+----+---------+-----+-------+---------------------+",
1356 "| id | name | age | score | score_regrintercept |",
1357 "+----+---------+-----+-------+---------------------+",
1358 "| 1 | Alice | 25 | 85.5 | -93.1354359925789 |",
1359 "| 2 | Bob | 30 | 92.0 | -93.1354359925789 |",
1360 "| 3 | | 15 | 78.5 | -93.1354359925789 |",
1361 "| 4 | Charlie | 40 | 95.0 | -93.1354359925789 |",
1362 "| 5 | Dave | 25 | 88.5 | -93.1354359925789 |",
1363 "+----+---------+-----+-------+---------------------+",
1364 ];
1365
1366 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1367 }
1368
1369 #[tokio::test]
1370 async fn test_regr_r2_rule() {
1371 let df = create_test_df().await;
1372 let rule = dfq_regr_r2(Some(col("age")), None);
1373 let result = rule.apply(df.clone(), "score").unwrap();
1374
1375 let expected = vec![
1376 "+----+---------+-----+-------+--------------------+",
1377 "| id | name | age | score | score_regrr2 |",
1378 "+----+---------+-----+-------+--------------------+",
1379 "| 1 | Alice | 25 | 85.5 | 0.9152939412679669 |",
1380 "| 2 | Bob | 30 | 92.0 | 0.9152939412679669 |",
1381 "| 3 | | 15 | 78.5 | 0.9152939412679669 |",
1382 "| 4 | Charlie | 40 | 95.0 | 0.9152939412679669 |",
1383 "| 5 | Dave | 25 | 88.5 | 0.9152939412679669 |",
1384 "+----+---------+-----+-------+--------------------+",
1385 ];
1386
1387 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1388
1389 let rule = dfq_regr_r2(None, Some(col("age")));
1390 let result = rule.apply(df, "score").unwrap();
1391
1392 let expected = vec![
1393 "+----+---------+-----+-------+--------------------+",
1394 "| id | name | age | score | score_regrr2 |",
1395 "+----+---------+-----+-------+--------------------+",
1396 "| 1 | Alice | 25 | 85.5 | 0.9152939412679678 |",
1397 "| 2 | Bob | 30 | 92.0 | 0.9152939412679678 |",
1398 "| 3 | | 15 | 78.5 | 0.9152939412679678 |",
1399 "| 4 | Charlie | 40 | 95.0 | 0.9152939412679678 |",
1400 "| 5 | Dave | 25 | 88.5 | 0.9152939412679678 |",
1401 "+----+---------+-----+-------+--------------------+",
1402 ];
1403
1404 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1405 }
1406
1407 #[tokio::test]
1408 async fn test_regr_slope_rule() {
1409 let df = create_test_df().await;
1410 let rule = dfq_regr_slope(Some(col("age")), None);
1411 let result = rule.apply(df.clone(), "score").unwrap();
1412
1413 let expected = vec![
1414 "+----+---------+-----+-------+--------------------+",
1415 "| id | name | age | score | score_regrslope |",
1416 "+----+---------+-----+-------+--------------------+",
1417 "| 1 | Alice | 25 | 85.5 | 0.6696969696969696 |",
1418 "| 2 | Bob | 30 | 92.0 | 0.6696969696969696 |",
1419 "| 3 | | 15 | 78.5 | 0.6696969696969696 |",
1420 "| 4 | Charlie | 40 | 95.0 | 0.6696969696969696 |",
1421 "| 5 | Dave | 25 | 88.5 | 0.6696969696969696 |",
1422 "+----+---------+-----+-------+--------------------+",
1423 ];
1424
1425 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1426
1427 let rule = dfq_regr_slope(None, Some(col("age")));
1428 let result = rule.apply(df, "score").unwrap();
1429
1430 let expected = vec![
1431 "+----+---------+-----+-------+-------------------+",
1432 "| id | name | age | score | score_regrslope |",
1433 "+----+---------+-----+-------+-------------------+",
1434 "| 1 | Alice | 25 | 85.5 | 1.366728509585653 |",
1435 "| 2 | Bob | 30 | 92.0 | 1.366728509585653 |",
1436 "| 3 | | 15 | 78.5 | 1.366728509585653 |",
1437 "| 4 | Charlie | 40 | 95.0 | 1.366728509585653 |",
1438 "| 5 | Dave | 25 | 88.5 | 1.366728509585653 |",
1439 "+----+---------+-----+-------+-------------------+",
1440 ];
1441
1442 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1443 }
1444
1445 #[tokio::test]
1446 async fn test_regr_sxx_rule() {
1447 let df = create_test_df().await;
1448 let rule = dfq_regr_sxx(Some(col("age")), None);
1449 let result = rule.apply(df.clone(), "score").unwrap();
1450
1451 let expected = vec![
1452 "+----+---------+-----+-------+---------------+",
1453 "| id | name | age | score | score_regrsxx |",
1454 "+----+---------+-----+-------+---------------+",
1455 "| 1 | Alice | 25 | 85.5 | 330.0 |",
1456 "| 2 | Bob | 30 | 92.0 | 330.0 |",
1457 "| 3 | | 15 | 78.5 | 330.0 |",
1458 "| 4 | Charlie | 40 | 95.0 | 330.0 |",
1459 "| 5 | Dave | 25 | 88.5 | 330.0 |",
1460 "+----+---------+-----+-------+---------------+",
1461 ];
1462
1463 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1464
1465 let rule = dfq_regr_sxx(None, Some(col("age")));
1466 let result = rule.apply(df, "score").unwrap();
1467
1468 let expected = vec![
1469 "+----+---------+-----+-------+---------------+",
1470 "| id | name | age | score | score_regrsxx |",
1471 "+----+---------+-----+-------+---------------+",
1472 "| 1 | Alice | 25 | 85.5 | 161.7 |",
1473 "| 2 | Bob | 30 | 92.0 | 161.7 |",
1474 "| 3 | | 15 | 78.5 | 161.7 |",
1475 "| 4 | Charlie | 40 | 95.0 | 161.7 |",
1476 "| 5 | Dave | 25 | 88.5 | 161.7 |",
1477 "+----+---------+-----+-------+---------------+",
1478 ];
1479
1480 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1481 }
1482
1483 #[tokio::test]
1484 async fn test_regr_sxy_rule() {
1485 let df = create_test_df().await;
1486 let rule = dfq_regr_sxy(Some(col("age")), None);
1487 let result = rule.apply(df.clone(), "score").unwrap();
1488
1489 let expected = vec![
1490 "+----+---------+-----+-------+--------------------+",
1491 "| id | name | age | score | score_regrsxy |",
1492 "+----+---------+-----+-------+--------------------+",
1493 "| 1 | Alice | 25 | 85.5 | 220.99999999999994 |",
1494 "| 2 | Bob | 30 | 92.0 | 220.99999999999994 |",
1495 "| 3 | | 15 | 78.5 | 220.99999999999994 |",
1496 "| 4 | Charlie | 40 | 95.0 | 220.99999999999994 |",
1497 "| 5 | Dave | 25 | 88.5 | 220.99999999999994 |",
1498 "+----+---------+-----+-------+--------------------+",
1499 ];
1500
1501 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1502
1503 let rule = dfq_regr_sxy(None, Some(col("age")));
1504 let result = rule.apply(df, "score").unwrap();
1505
1506 let expected = vec![
1507 "+----+---------+-----+-------+--------------------+",
1508 "| id | name | age | score | score_regrsxy |",
1509 "+----+---------+-----+-------+--------------------+",
1510 "| 1 | Alice | 25 | 85.5 | 221.00000000000006 |",
1511 "| 2 | Bob | 30 | 92.0 | 221.00000000000006 |",
1512 "| 3 | | 15 | 78.5 | 221.00000000000006 |",
1513 "| 4 | Charlie | 40 | 95.0 | 221.00000000000006 |",
1514 "| 5 | Dave | 25 | 88.5 | 221.00000000000006 |",
1515 "+----+---------+-----+-------+--------------------+",
1516 ];
1517
1518 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1519 }
1520
1521 #[tokio::test]
1522 async fn test_regr_syy_rule() {
1523 let df = create_test_df().await;
1524 let rule = dfq_regr_syy(Some(col("age")), None);
1525 let result = rule.apply(df.clone(), "score").unwrap();
1526
1527 let expected = vec![
1528 "+----+---------+-----+-------+---------------+",
1529 "| id | name | age | score | score_regrsyy |",
1530 "+----+---------+-----+-------+---------------+",
1531 "| 1 | Alice | 25 | 85.5 | 161.7 |",
1532 "| 2 | Bob | 30 | 92.0 | 161.7 |",
1533 "| 3 | | 15 | 78.5 | 161.7 |",
1534 "| 4 | Charlie | 40 | 95.0 | 161.7 |",
1535 "| 5 | Dave | 25 | 88.5 | 161.7 |",
1536 "+----+---------+-----+-------+---------------+",
1537 ];
1538
1539 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1540
1541 let rule = dfq_regr_syy(None, Some(col("age")));
1542 let result = rule.apply(df, "score").unwrap();
1543
1544 let expected = vec![
1545 "+----+---------+-----+-------+---------------+",
1546 "| id | name | age | score | score_regrsyy |",
1547 "+----+---------+-----+-------+---------------+",
1548 "| 1 | Alice | 25 | 85.5 | 330.0 |",
1549 "| 2 | Bob | 30 | 92.0 | 330.0 |",
1550 "| 3 | | 15 | 78.5 | 330.0 |",
1551 "| 4 | Charlie | 40 | 95.0 | 330.0 |",
1552 "| 5 | Dave | 25 | 88.5 | 330.0 |",
1553 "+----+---------+-----+-------+---------------+",
1554 ];
1555
1556 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1557 }
1558
1559 #[tokio::test]
1560 async fn test_nth_value_rule() {
1561 let df = create_test_df().await;
1562 let rule = dfq_nth_value(2, None);
1563 let result = rule.apply(df.clone(), "score").unwrap();
1564
1565 let expected = vec![
1566 "+----+---------+-----+-------+----------------+",
1567 "| id | name | age | score | score_nthvalue |",
1568 "+----+---------+-----+-------+----------------+",
1569 "| 1 | Alice | 25 | 85.5 | 92.0 |",
1570 "| 2 | Bob | 30 | 92.0 | 92.0 |",
1571 "| 3 | | 15 | 78.5 | 92.0 |",
1572 "| 4 | Charlie | 40 | 95.0 | 92.0 |",
1573 "| 5 | Dave | 25 | 88.5 | 92.0 |",
1574 "+----+---------+-----+-------+----------------+",
1575 ];
1576
1577 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1578
1579 let rule = dfq_nth_value(2, Some(vec![col("score").sort(true, false)]));
1580 let result = rule.apply(df, "score").unwrap();
1581
1582 let expected = vec![
1583 "+----+---------+-----+-------+----------------+",
1584 "| id | name | age | score | score_nthvalue |",
1585 "+----+---------+-----+-------+----------------+",
1586 "| 1 | Alice | 25 | 85.5 | 85.5 |",
1587 "| 2 | Bob | 30 | 92.0 | 85.5 |",
1588 "| 3 | | 15 | 78.5 | 85.5 |",
1589 "| 4 | Charlie | 40 | 95.0 | 85.5 |",
1590 "| 5 | Dave | 25 | 88.5 | 85.5 |",
1591 "+----+---------+-----+-------+----------------+",
1592 ];
1593
1594 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1595 }
1596
1597 #[tokio::test]
1598 async fn test_first_value_rule() {
1599 let df = create_test_df().await;
1600 let rule = dfq_first_value(None);
1601 let result = rule.apply(df.clone(), "score").unwrap();
1602
1603 let expected = vec![
1604 "+----+---------+-----+-------+------------------+",
1605 "| id | name | age | score | score_firstvalue |",
1606 "+----+---------+-----+-------+------------------+",
1607 "| 1 | Alice | 25 | 85.5 | 85.5 |",
1608 "| 2 | Bob | 30 | 92.0 | 85.5 |",
1609 "| 3 | | 15 | 78.5 | 85.5 |",
1610 "| 4 | Charlie | 40 | 95.0 | 85.5 |",
1611 "| 5 | Dave | 25 | 88.5 | 85.5 |",
1612 "+----+---------+-----+-------+------------------+",
1613 ];
1614
1615 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1616
1617 let rule = dfq_first_value(Some(vec![col("score").sort(true, false)]));
1618 let result = rule.apply(df, "score").unwrap();
1619
1620 let expected = vec![
1621 "+----+---------+-----+-------+------------------+",
1622 "| id | name | age | score | score_firstvalue |",
1623 "+----+---------+-----+-------+------------------+",
1624 "| 1 | Alice | 25 | 85.5 | 78.5 |",
1625 "| 2 | Bob | 30 | 92.0 | 78.5 |",
1626 "| 3 | | 15 | 78.5 | 78.5 |",
1627 "| 4 | Charlie | 40 | 95.0 | 78.5 |",
1628 "| 5 | Dave | 25 | 88.5 | 78.5 |",
1629 "+----+---------+-----+-------+------------------+",
1630 ];
1631
1632 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1633 }
1634
1635 #[tokio::test]
1636 async fn test_custom_aggregation_rule_builder() {
1637 let df = create_test_df().await;
1638
1639 let group_by_rule =
1641 CustomAggregationRule::builder(col("max_score"), "max_score_by_name".to_string())
1642 .with_group_by(vec![col("name")])
1643 .with_aggregate_exprs(vec![max(col("score")).alias("max_score")])
1644 .build();
1645
1646 let group_by_result = group_by_rule
1647 .apply(df.clone(), "score")
1648 .unwrap()
1649 .sort(vec![col("id").sort(true, false)])
1650 .unwrap();
1651
1652 let group_by_expected = vec![
1653 "+----+---------+-----+-------+-------------------------+",
1654 "| id | name | age | score | score_max_score_by_name |",
1655 "+----+---------+-----+-------+-------------------------+",
1656 "| 1 | Alice | 25 | 85.5 | 92.0 |",
1657 "| 1 | Alice | 25 | 85.5 | 85.5 |",
1658 "| 1 | Alice | 25 | 85.5 | 78.5 |",
1659 "| 1 | Alice | 25 | 85.5 | 88.5 |",
1660 "| 1 | Alice | 25 | 85.5 | 95.0 |",
1661 "| 2 | Bob | 30 | 92.0 | 92.0 |",
1662 "| 2 | Bob | 30 | 92.0 | 85.5 |",
1663 "| 2 | Bob | 30 | 92.0 | 78.5 |",
1664 "| 2 | Bob | 30 | 92.0 | 88.5 |",
1665 "| 2 | Bob | 30 | 92.0 | 95.0 |",
1666 "| 3 | | 15 | 78.5 | 92.0 |",
1667 "| 3 | | 15 | 78.5 | 85.5 |",
1668 "| 3 | | 15 | 78.5 | 78.5 |",
1669 "| 3 | | 15 | 78.5 | 88.5 |",
1670 "| 3 | | 15 | 78.5 | 95.0 |",
1671 "| 4 | Charlie | 40 | 95.0 | 92.0 |",
1672 "| 4 | Charlie | 40 | 95.0 | 85.5 |",
1673 "| 4 | Charlie | 40 | 95.0 | 78.5 |",
1674 "| 4 | Charlie | 40 | 95.0 | 88.5 |",
1675 "| 4 | Charlie | 40 | 95.0 | 95.0 |",
1676 "| 5 | Dave | 25 | 88.5 | 92.0 |",
1677 "| 5 | Dave | 25 | 88.5 | 85.5 |",
1678 "| 5 | Dave | 25 | 88.5 | 78.5 |",
1679 "| 5 | Dave | 25 | 88.5 | 88.5 |",
1680 "| 5 | Dave | 25 | 88.5 | 95.0 |",
1681 "+----+---------+-----+-------+-------------------------+",
1682 ];
1683
1684 assert_batches_eq!(
1685 &group_by_expected,
1686 &group_by_result.collect().await.unwrap()
1687 );
1688
1689 let filter_rule = CustomAggregationRule::builder(
1691 col("max_score_non_null"),
1692 "max_score_non_null".to_string(),
1693 )
1694 .with_filter(col("name").is_not_null())
1695 .with_aggregate_exprs(vec![max(col("score")).alias("max_score_non_null")])
1696 .build();
1697
1698 let filter_result = filter_rule.apply(df.clone(), "score").unwrap();
1699
1700 let filter_expected = vec![
1701 "+----+---------+-----+-------+--------------------------+",
1702 "| id | name | age | score | score_max_score_non_null |",
1703 "+----+---------+-----+-------+--------------------------+",
1704 "| 1 | Alice | 25 | 85.5 | 95.0 |",
1705 "| 2 | Bob | 30 | 92.0 | 95.0 |",
1706 "| 3 | | 15 | 78.5 | 95.0 |",
1707 "| 4 | Charlie | 40 | 95.0 | 95.0 |",
1708 "| 5 | Dave | 25 | 88.5 | 95.0 |",
1709 "+----+---------+-----+-------+--------------------------+",
1710 ];
1711
1712 assert_batches_eq!(&filter_expected, &filter_result.collect().await.unwrap());
1713 }
1714}