Skip to main content

datafusion_quality/rules/
table.rs

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/// Rule that counts null values in a column across the entire table
10#[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 to create a zero argument calculation rule
307#[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
638/// Returns the nth value in the column.
639///
640/// # Arguments
641///
642/// * `n` - The nth value to return (1-based index)
643/// * `sort_exprs` - Optional sort expressions to determine the order
644///
645/// # Examples
646///
647/// ```
648/// use datafusion_quality::rules::table::dfq_nth_value;
649/// use datafusion_quality::RuleSet;
650///
651/// // Create a rule to get the 3rd value in the age column
652/// let rule = dfq_nth_value(3, None);
653/// let mut ruleset = RuleSet::new();
654/// ruleset.with_table_rule("age", rule, None);
655/// ```
656pub 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
662/// Returns the first value in the column.
663///
664/// # Arguments
665///
666/// * `sort_exprs` - Optional sort expressions to determine the order
667///
668/// # Examples
669///
670/// ```
671/// use datafusion_quality::rules::table::dfq_first_value;
672/// use datafusion_quality::RuleSet;
673///
674/// // Create a rule to get the first value in the age column
675/// let rule = dfq_first_value(None);
676/// let mut ruleset = RuleSet::new();
677/// ruleset.with_table_rule("age", rule, None);
678/// ```
679pub 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/// Rule that applies a custom aggregation across the entire table
744#[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        // Test negated null count rule
885        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        // Test with_group_by - groups by name and gets max score per group
1640        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        // Test with_filter - filters out null names and gets max score
1690        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}