1use crate::{ColumnRule, ValidationError, error::DataFusionSnafu};
2use datafusion::{logical_expr::Between, prelude::*};
3use snafu::ResultExt;
4use std::sync::Arc;
5
6#[derive(Debug, Clone, Default)]
8pub struct NullRule {
9 negated: Option<bool>,
10}
11
12impl NullRule {
13 pub fn new(negated: Option<bool>) -> Self {
14 Self { negated }
15 }
16}
17
18impl ColumnRule for NullRule {
19 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
20 let col = col(column_name);
21 let is_not_null = if self.negated.unwrap_or_default() {
22 col.is_not_null()
23 } else {
24 col.is_null()
25 };
26
27 df.with_column(&self.new_column_name(column_name), is_not_null)
28 .context(DataFusionSnafu)
29 }
30
31 fn name(&self) -> &str {
32 if self.negated.unwrap_or_default() {
33 "not_null"
34 } else {
35 "null"
36 }
37 }
38
39 fn new_column_name(&self, column_name: &str) -> String {
40 format!("{}_{}", column_name, self.name())
41 }
42
43 fn description(&self) -> &str {
44 "Checks if values in a column are null/not null"
45 }
46}
47
48pub fn dfq_not_null() -> Arc<NullRule> {
66 Arc::new(NullRule::new(Some(true)))
67}
68
69pub fn dfq_null() -> Arc<NullRule> {
87 Arc::new(NullRule::new(Some(false)))
88}
89
90#[derive(Debug, Clone)]
92pub struct RangeRule {
93 min: f64,
94 max: f64,
95 negated: Option<bool>,
96}
97
98impl RangeRule {
99 pub fn new(min: f64, max: f64, negated: Option<bool>) -> Self {
100 Self { min, max, negated }
101 }
102}
103
104impl ColumnRule for RangeRule {
105 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
106 let col = col(column_name);
107 let in_range = Expr::Between(Between {
108 expr: Box::new(col),
109 negated: self.negated.unwrap_or(false),
110 low: Box::new(lit(self.min)),
111 high: Box::new(lit(self.max)),
112 });
113
114 df.with_column(&self.new_column_name(column_name), in_range)
115 .context(DataFusionSnafu)
116 }
117
118 fn name(&self) -> &str {
119 if self.negated.unwrap_or_default() {
120 "not_in_range"
121 } else {
122 "in_range"
123 }
124 }
125
126 fn new_column_name(&self, column_name: &str) -> String {
127 format!("{}_{}", column_name, self.name())
128 }
129
130 fn description(&self) -> &str {
131 "Checks if values in a column (does not) fall within a specified range"
132 }
133}
134
135pub fn dfq_in_range(min: f64, max: f64) -> Arc<RangeRule> {
154 Arc::new(RangeRule::new(min, max, None))
155}
156
157pub fn dfq_not_in_range(min: f64, max: f64) -> Arc<RangeRule> {
176 Arc::new(RangeRule::new(min, max, Some(true)))
177}
178
179#[derive(Debug, Clone)]
181pub struct PatternRule {
182 pattern: String,
183 negated: Option<bool>,
184 case_sensitive: Option<bool>,
185}
186
187impl PatternRule {
188 pub fn new(pattern: &str, negated: Option<bool>, case_sensitive: Option<bool>) -> Self {
189 Self {
190 pattern: pattern.to_string(),
191 negated,
192 case_sensitive,
193 }
194 }
195}
196
197impl ColumnRule for PatternRule {
198 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
199 let col = col(column_name);
200 let matches_pattern = match (
201 self.negated.unwrap_or_default(),
202 self.case_sensitive.unwrap_or_default(),
203 ) {
204 (true, true) => col.not_like(lit(&self.pattern)),
205 (false, true) => col.like(lit(&self.pattern)),
206 (true, false) => col.not_ilike(lit(&self.pattern)),
207 (false, false) => col.ilike(lit(&self.pattern)),
208 };
209
210 df.with_column(&self.new_column_name(column_name), matches_pattern)
211 .context(DataFusionSnafu)
212 }
213
214 fn name(&self) -> &str {
215 match (
216 self.negated.unwrap_or_default(),
217 self.case_sensitive.unwrap_or_default(),
218 ) {
219 (true, true) => "not_like",
220 (false, true) => "like",
221 (true, false) => "not_ilike",
222 (false, false) => "ilike",
223 }
224 }
225
226 fn new_column_name(&self, column_name: &str) -> String {
227 format!("{}_{}", column_name, self.name())
228 }
229
230 fn description(&self) -> &str {
231 "Checks if values in a column match a pattern"
232 }
233}
234
235pub fn dfq_like(pattern: &str) -> Arc<PatternRule> {
253 Arc::new(PatternRule::new(pattern, Some(false), Some(true)))
254}
255
256pub fn dfq_not_like(pattern: &str) -> Arc<PatternRule> {
274 Arc::new(PatternRule::new(pattern, Some(true), Some(true)))
275}
276
277pub fn dfq_ilike(pattern: &str) -> Arc<PatternRule> {
295 Arc::new(PatternRule::new(pattern, Some(false), Some(false)))
296}
297
298pub fn dfq_not_ilike(pattern: &str) -> Arc<PatternRule> {
316 Arc::new(PatternRule::new(pattern, Some(true), Some(false)))
317}
318
319#[derive(Debug, Clone)]
321pub struct ComparisonRule {
322 value: Expr,
323 negated: bool,
324 equals: bool,
325 comparison_type: ComparisonType,
326}
327
328#[derive(Debug, Clone, Copy)]
329pub enum ComparisonType {
330 LessThan,
331 GreaterThan,
332 Equals,
333}
334
335impl ComparisonRule {
336 pub fn new(value: Expr, negated: bool, equals: bool, comparison_type: ComparisonType) -> Self {
337 Self {
338 value,
339 negated,
340 equals,
341 comparison_type,
342 }
343 }
344}
345
346impl ColumnRule for ComparisonRule {
347 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
348 let col = col(column_name);
349 let comparison = match (self.comparison_type, self.equals) {
350 (ComparisonType::LessThan, true) => col.lt_eq(self.value.clone()),
351 (ComparisonType::LessThan, false) => col.lt(self.value.clone()),
352 (ComparisonType::GreaterThan, true) => col.gt_eq(self.value.clone()),
353 (ComparisonType::GreaterThan, false) => col.gt(self.value.clone()),
354 (ComparisonType::Equals, _) => col.eq(self.value.clone()),
355 };
356
357 let expr = if self.negated {
358 comparison.not()
359 } else {
360 comparison
361 };
362
363 df.with_column(&self.new_column_name(column_name), expr)
364 .context(DataFusionSnafu)
365 }
366
367 fn name(&self) -> &str {
368 match (self.comparison_type, self.negated, self.equals) {
369 (ComparisonType::LessThan, false, false) => "less_than",
370 (ComparisonType::LessThan, false, true) => "less_than_equals",
371 (ComparisonType::LessThan, true, false) => "not_less_than",
372 (ComparisonType::LessThan, true, true) => "not_less_than_equals",
373 (ComparisonType::GreaterThan, false, false) => "greater_than",
374 (ComparisonType::GreaterThan, false, true) => "greater_than_equals",
375 (ComparisonType::GreaterThan, true, false) => "not_greater_than",
376 (ComparisonType::GreaterThan, true, true) => "not_greater_than_equals",
377 (ComparisonType::Equals, false, _) => "equals",
378 (ComparisonType::Equals, true, _) => "not_equals",
379 }
380 }
381
382 fn new_column_name(&self, column_name: &str) -> String {
383 format!("{}_{}", column_name, self.name())
384 }
385
386 fn description(&self) -> &str {
387 "Checks if values in a column satisfy a comparison with a value"
388 }
389}
390
391pub fn dfq_lt(value: Expr) -> Arc<ComparisonRule> {
410 Arc::new(ComparisonRule::new(
411 value,
412 false,
413 false,
414 ComparisonType::LessThan,
415 ))
416}
417
418pub fn dfq_lte(value: Expr) -> Arc<ComparisonRule> {
437 Arc::new(ComparisonRule::new(
438 value,
439 false,
440 true,
441 ComparisonType::LessThan,
442 ))
443}
444
445pub fn dfq_not_lt(value: Expr) -> Arc<ComparisonRule> {
464 Arc::new(ComparisonRule::new(
465 value,
466 true,
467 false,
468 ComparisonType::LessThan,
469 ))
470}
471
472pub fn dfq_not_lte(value: Expr) -> Arc<ComparisonRule> {
491 Arc::new(ComparisonRule::new(
492 value,
493 true,
494 true,
495 ComparisonType::LessThan,
496 ))
497}
498
499pub fn dfq_gt(value: Expr) -> Arc<ComparisonRule> {
518 Arc::new(ComparisonRule::new(
519 value,
520 false,
521 false,
522 ComparisonType::GreaterThan,
523 ))
524}
525
526pub fn dfq_gte(value: Expr) -> Arc<ComparisonRule> {
545 Arc::new(ComparisonRule::new(
546 value,
547 false,
548 true,
549 ComparisonType::GreaterThan,
550 ))
551}
552
553pub fn dfq_not_gt(value: Expr) -> Arc<ComparisonRule> {
572 Arc::new(ComparisonRule::new(
573 value,
574 true,
575 false,
576 ComparisonType::GreaterThan,
577 ))
578}
579
580pub fn dfq_not_gte(value: Expr) -> Arc<ComparisonRule> {
599 Arc::new(ComparisonRule::new(
600 value,
601 true,
602 true,
603 ComparisonType::GreaterThan,
604 ))
605}
606
607pub fn dfq_eq(value: Expr) -> Arc<ComparisonRule> {
626 Arc::new(ComparisonRule::new(
627 value,
628 false,
629 false,
630 ComparisonType::Equals,
631 ))
632}
633
634pub fn dfq_not_eq(value: Expr) -> Arc<ComparisonRule> {
653 Arc::new(ComparisonRule::new(
654 value,
655 true,
656 false,
657 ComparisonType::Equals,
658 ))
659}
660
661#[derive(Debug, Clone)]
662pub struct LengthRule {
663 min: Option<u32>,
664 max: Option<u32>,
665}
666
667impl LengthRule {
668 pub fn new(min: Option<u32>, max: Option<u32>) -> Self {
669 Self { min, max }
670 }
671}
672
673impl ColumnRule for LengthRule {
674 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
675 let mut expr = char_length(col(column_name));
676
677 match (self.min, self.max) {
678 (Some(min), Some(max)) => {
679 expr = expr.between(lit(min), lit(max));
680 }
681 (Some(min), None) => {
682 expr = expr.gt_eq(lit(min));
683 }
684 (None, Some(max)) => {
685 expr = expr.lt(lit(max));
686 }
687 (None, None) => {
688 return Err(ValidationError::Configuration {
689 message: "Length rule must have either a minimum or maximum length".to_string(),
690 });
691 }
692 }
693
694 df.with_column(&self.new_column_name(column_name), expr)
695 .context(DataFusionSnafu)
696 }
697
698 fn name(&self) -> &str {
699 match (self.min, self.max) {
700 (Some(_), Some(_)) => "length_range",
701 (Some(_), None) => "min_length",
702 (None, Some(_)) => "max_length",
703 (None, None) => "length",
704 }
705 }
706
707 fn new_column_name(&self, column_name: &str) -> String {
708 format!("{}_length", column_name)
709 }
710
711 fn description(&self) -> &str {
712 "Checks if the length of a column is between a minimum and maximum value"
713 }
714}
715
716pub fn dfq_str_length(min: Option<u32>, max: Option<u32>) -> Arc<LengthRule> {
735 Arc::new(LengthRule::new(min, max))
736}
737
738pub fn dfq_str_min_length(min: u32) -> Arc<LengthRule> {
756 Arc::new(LengthRule::new(Some(min), None))
757}
758
759pub fn dfq_str_max_length(max: u32) -> Arc<LengthRule> {
777 Arc::new(LengthRule::new(None, Some(max)))
778}
779
780pub fn dfq_str_empty() -> Arc<LengthRule> {
794 Arc::new(LengthRule::new(None, Some(0)))
795}
796
797pub fn dfq_str_not_empty() -> Arc<LengthRule> {
811 Arc::new(LengthRule::new(Some(1), None))
812}
813
814#[derive(Debug, Clone)]
816pub struct CustomRule {
817 rule_name: String,
818 expression: Expr,
819}
820
821impl CustomRule {
822 pub fn new(rule_name: &str, expression: Expr) -> Self {
823 Self {
824 rule_name: rule_name.to_string(),
825 expression,
826 }
827 }
828}
829
830impl ColumnRule for CustomRule {
831 fn apply(&self, df: DataFrame, column_name: &str) -> Result<DataFrame, ValidationError> {
832 let expr = self.expression.clone();
833 df.with_column(&self.new_column_name(column_name), expr)
834 .context(DataFusionSnafu)
835 }
836
837 fn name(&self) -> &str {
838 "custom"
839 }
840
841 fn new_column_name(&self, column_name: &str) -> String {
842 format!("{}_{}", column_name, self.rule_name)
843 }
844
845 fn description(&self) -> &str {
846 "Applies a custom SQL expression to a column"
847 }
848}
849
850pub fn dfq_custom(rule_name: &str, expression: Expr) -> Arc<CustomRule> {
870 Arc::new(CustomRule::new(rule_name, expression))
871}
872
873#[cfg(test)]
874mod tests {
875 use super::*;
876 use arrow::array::{Float64Array, Int32Array, StringArray};
877 use arrow::datatypes::{DataType, Field, Schema};
878 use arrow::record_batch::RecordBatch;
879 use datafusion::assert_batches_eq;
880
881 async fn create_test_df() -> DataFrame {
882 let schema = Schema::new(vec![
883 Field::new("id", DataType::Int32, false),
884 Field::new("name", DataType::Utf8, false),
885 Field::new("age", DataType::Int32, true),
886 Field::new("score", DataType::Float64, true),
887 ]);
888
889 let batch = RecordBatch::try_new(
890 Arc::new(schema),
891 vec![
892 Arc::new(Int32Array::from(vec![1, 2, 3])),
893 Arc::new(StringArray::from(vec!["Alice", "Bob", "Charlie"])),
894 Arc::new(Int32Array::from(vec![Some(25), None, Some(30)])),
895 Arc::new(Float64Array::from(vec![Some(85.5), Some(92.0), None])),
896 ],
897 )
898 .unwrap();
899
900 let ctx = SessionContext::new();
901 ctx.read_batch(batch).unwrap()
902 }
903
904 #[tokio::test]
905 async fn test_not_null_rule() {
906 let df = create_test_df().await;
907 let rule = dfq_null();
908 let result = rule.apply(df.clone(), "age").unwrap();
909
910 let expected = vec![
911 "+----+---------+-----+-------+----------+",
912 "| id | name | age | score | age_null |",
913 "+----+---------+-----+-------+----------+",
914 "| 1 | Alice | 25 | 85.5 | false |",
915 "| 2 | Bob | | 92.0 | true |",
916 "| 3 | Charlie | 30 | | false |",
917 "+----+---------+-----+-------+----------+",
918 ];
919
920 assert_batches_eq!(&expected, &result.collect().await.unwrap());
921
922 let df = create_test_df().await;
924 let rule = dfq_not_null();
925 let result = rule.apply(df, "age").unwrap();
926
927 let expected = vec![
928 "+----+---------+-----+-------+--------------+",
929 "| id | name | age | score | age_not_null |",
930 "+----+---------+-----+-------+--------------+",
931 "| 1 | Alice | 25 | 85.5 | true |",
932 "| 2 | Bob | | 92.0 | false |",
933 "| 3 | Charlie | 30 | | true |",
934 "+----+---------+-----+-------+--------------+",
935 ];
936
937 assert_batches_eq!(&expected, &result.collect().await.unwrap());
938 }
939
940 #[tokio::test]
941 async fn test_range_rule() {
942 let df = create_test_df().await;
943 let rule = dfq_in_range(0.0, 100.0);
944 let result = rule.apply(df.clone(), "score").unwrap();
945
946 let expected = vec![
947 "+----+---------+-----+-------+----------------+",
948 "| id | name | age | score | score_in_range |",
949 "+----+---------+-----+-------+----------------+",
950 "| 1 | Alice | 25 | 85.5 | true |",
951 "| 2 | Bob | | 92.0 | true |",
952 "| 3 | Charlie | 30 | | |",
953 "+----+---------+-----+-------+----------------+",
954 ];
955
956 assert_batches_eq!(&expected, &result.collect().await.unwrap());
957
958 let df = create_test_df().await;
960 let rule = dfq_not_in_range(0.0, 100.0);
961 let result = rule.apply(df, "score").unwrap();
962
963 let expected = vec![
964 "+----+---------+-----+-------+--------------------+",
965 "| id | name | age | score | score_not_in_range |",
966 "+----+---------+-----+-------+--------------------+",
967 "| 1 | Alice | 25 | 85.5 | false |",
968 "| 2 | Bob | | 92.0 | false |",
969 "| 3 | Charlie | 30 | | |",
970 "+----+---------+-----+-------+--------------------+",
971 ];
972
973 assert_batches_eq!(&expected, &result.collect().await.unwrap());
974 }
975
976 #[tokio::test]
977 async fn test_pattern_rule() {
978 let df = create_test_df().await;
980 let rule = dfq_like("A%");
981 let result = rule.apply(df, "name").unwrap();
982
983 let expected = vec![
984 "+----+---------+-----+-------+-----------+",
985 "| id | name | age | score | name_like |",
986 "+----+---------+-----+-------+-----------+",
987 "| 1 | Alice | 25 | 85.5 | true |",
988 "| 2 | Bob | | 92.0 | false |",
989 "| 3 | Charlie | 30 | | false |",
990 "+----+---------+-----+-------+-----------+",
991 ];
992
993 assert_batches_eq!(&expected, &result.collect().await.unwrap());
994
995 let df = create_test_df().await;
997 let rule = dfq_ilike("a%");
998 let result = rule.apply(df, "name").unwrap();
999
1000 let expected = vec![
1001 "+----+---------+-----+-------+------------+",
1002 "| id | name | age | score | name_ilike |",
1003 "+----+---------+-----+-------+------------+",
1004 "| 1 | Alice | 25 | 85.5 | true |",
1005 "| 2 | Bob | | 92.0 | false |",
1006 "| 3 | Charlie | 30 | | false |",
1007 "+----+---------+-----+-------+------------+",
1008 ];
1009
1010 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1011
1012 let df = create_test_df().await;
1014 let rule = dfq_not_like("A%");
1015 let result = rule.apply(df, "name").unwrap();
1016
1017 let expected = vec![
1018 "+----+---------+-----+-------+---------------+",
1019 "| id | name | age | score | name_not_like |",
1020 "+----+---------+-----+-------+---------------+",
1021 "| 1 | Alice | 25 | 85.5 | false |",
1022 "| 2 | Bob | | 92.0 | true |",
1023 "| 3 | Charlie | 30 | | true |",
1024 "+----+---------+-----+-------+---------------+",
1025 ];
1026
1027 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1028
1029 let df = create_test_df().await;
1031 let rule = dfq_not_ilike("a%");
1032 let result = rule.apply(df, "name").unwrap();
1033
1034 let expected = vec![
1035 "+----+---------+-----+-------+----------------+",
1036 "| id | name | age | score | name_not_ilike |",
1037 "+----+---------+-----+-------+----------------+",
1038 "| 1 | Alice | 25 | 85.5 | false |",
1039 "| 2 | Bob | | 92.0 | true |",
1040 "| 3 | Charlie | 30 | | true |",
1041 "+----+---------+-----+-------+----------------+",
1042 ];
1043
1044 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1045 }
1046
1047 #[tokio::test]
1048 async fn test_custom_rule() {
1049 let df = create_test_df().await;
1050 let rule = dfq_custom("age_gt_25", col("age").gt(lit(25)));
1051 let result = rule.apply(df, "age").unwrap();
1052
1053 let expected = vec![
1054 "+----+---------+-----+-------+---------------+",
1055 "| id | name | age | score | age_age_gt_25 |",
1056 "+----+---------+-----+-------+---------------+",
1057 "| 1 | Alice | 25 | 85.5 | false |",
1058 "| 2 | Bob | | 92.0 | |",
1059 "| 3 | Charlie | 30 | | true |",
1060 "+----+---------+-----+-------+---------------+",
1061 ];
1062
1063 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1064 }
1065
1066 #[tokio::test]
1067 async fn test_less_than_rule() {
1068 let df = create_test_df().await;
1069 let rule = dfq_lt(lit(30));
1070 let result = rule.apply(df, "age").unwrap();
1071
1072 let expected = vec![
1073 "+----+---------+-----+-------+---------------+",
1074 "| id | name | age | score | age_less_than |",
1075 "+----+---------+-----+-------+---------------+",
1076 "| 1 | Alice | 25 | 85.5 | true |",
1077 "| 2 | Bob | | 92.0 | |",
1078 "| 3 | Charlie | 30 | | false |",
1079 "+----+---------+-----+-------+---------------+",
1080 ];
1081
1082 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1083 }
1084
1085 #[tokio::test]
1086 async fn test_less_than_equals_rule() {
1087 let df = create_test_df().await;
1088 let rule = dfq_lte(lit(30));
1089 let result = rule.apply(df, "age").unwrap();
1090
1091 let expected = vec![
1092 "+----+---------+-----+-------+----------------------+",
1093 "| id | name | age | score | age_less_than_equals |",
1094 "+----+---------+-----+-------+----------------------+",
1095 "| 1 | Alice | 25 | 85.5 | true |",
1096 "| 2 | Bob | | 92.0 | |",
1097 "| 3 | Charlie | 30 | | true |",
1098 "+----+---------+-----+-------+----------------------+",
1099 ];
1100
1101 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1102 }
1103
1104 #[tokio::test]
1105 async fn test_not_less_than_rule() {
1106 let df = create_test_df().await;
1107 let rule = dfq_not_lt(lit(30));
1108 let result = rule.apply(df, "age").unwrap();
1109
1110 let expected = vec![
1111 "+----+---------+-----+-------+-------------------+",
1112 "| id | name | age | score | age_not_less_than |",
1113 "+----+---------+-----+-------+-------------------+",
1114 "| 1 | Alice | 25 | 85.5 | false |",
1115 "| 2 | Bob | | 92.0 | |",
1116 "| 3 | Charlie | 30 | | true |",
1117 "+----+---------+-----+-------+-------------------+",
1118 ];
1119
1120 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1121 }
1122
1123 #[tokio::test]
1124 async fn test_not_less_than_equals_rule() {
1125 let df = create_test_df().await;
1126 let rule = dfq_not_lte(lit(30));
1127 let result = rule.apply(df, "age").unwrap();
1128
1129 let expected = vec![
1130 "+----+---------+-----+-------+--------------------------+",
1131 "| id | name | age | score | age_not_less_than_equals |",
1132 "+----+---------+-----+-------+--------------------------+",
1133 "| 1 | Alice | 25 | 85.5 | false |",
1134 "| 2 | Bob | | 92.0 | |",
1135 "| 3 | Charlie | 30 | | false |",
1136 "+----+---------+-----+-------+--------------------------+",
1137 ];
1138
1139 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1140 }
1141
1142 #[tokio::test]
1143 async fn test_greater_than_rule() {
1144 let df = create_test_df().await;
1145 let rule = dfq_gt(lit(25));
1146 let result = rule.apply(df, "age").unwrap();
1147
1148 let expected = vec![
1149 "+----+---------+-----+-------+------------------+",
1150 "| id | name | age | score | age_greater_than |",
1151 "+----+---------+-----+-------+------------------+",
1152 "| 1 | Alice | 25 | 85.5 | false |",
1153 "| 2 | Bob | | 92.0 | |",
1154 "| 3 | Charlie | 30 | | true |",
1155 "+----+---------+-----+-------+------------------+",
1156 ];
1157
1158 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1159 }
1160
1161 #[tokio::test]
1162 async fn test_greater_than_equals_rule() {
1163 let df = create_test_df().await;
1164 let rule = dfq_gte(lit(25));
1165 let result = rule.apply(df, "age").unwrap();
1166
1167 let expected = vec![
1168 "+----+---------+-----+-------+-------------------------+",
1169 "| id | name | age | score | age_greater_than_equals |",
1170 "+----+---------+-----+-------+-------------------------+",
1171 "| 1 | Alice | 25 | 85.5 | true |",
1172 "| 2 | Bob | | 92.0 | |",
1173 "| 3 | Charlie | 30 | | true |",
1174 "+----+---------+-----+-------+-------------------------+",
1175 ];
1176
1177 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1178 }
1179
1180 #[tokio::test]
1181 async fn test_not_greater_than_rule() {
1182 let df = create_test_df().await;
1183 let rule = dfq_not_gt(lit(25));
1184 let result = rule.apply(df, "age").unwrap();
1185
1186 let expected = vec![
1187 "+----+---------+-----+-------+----------------------+",
1188 "| id | name | age | score | age_not_greater_than |",
1189 "+----+---------+-----+-------+----------------------+",
1190 "| 1 | Alice | 25 | 85.5 | true |",
1191 "| 2 | Bob | | 92.0 | |",
1192 "| 3 | Charlie | 30 | | false |",
1193 "+----+---------+-----+-------+----------------------+",
1194 ];
1195
1196 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1197 }
1198
1199 #[tokio::test]
1200 async fn test_not_greater_than_equals_rule() {
1201 let df = create_test_df().await;
1202 let rule = dfq_not_gte(lit(25));
1203 let result = rule.apply(df, "age").unwrap();
1204
1205 let expected = vec![
1206 "+----+---------+-----+-------+-----------------------------+",
1207 "| id | name | age | score | age_not_greater_than_equals |",
1208 "+----+---------+-----+-------+-----------------------------+",
1209 "| 1 | Alice | 25 | 85.5 | false |",
1210 "| 2 | Bob | | 92.0 | |",
1211 "| 3 | Charlie | 30 | | false |",
1212 "+----+---------+-----+-------+-----------------------------+",
1213 ];
1214
1215 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1216 }
1217
1218 #[tokio::test]
1219 async fn test_equals_rule() {
1220 let df = create_test_df().await;
1221 let rule = dfq_eq(lit(25));
1222 let result = rule.apply(df, "age").unwrap();
1223
1224 let expected = vec![
1225 "+----+---------+-----+-------+------------+",
1226 "| id | name | age | score | age_equals |",
1227 "+----+---------+-----+-------+------------+",
1228 "| 1 | Alice | 25 | 85.5 | true |",
1229 "| 2 | Bob | | 92.0 | |",
1230 "| 3 | Charlie | 30 | | false |",
1231 "+----+---------+-----+-------+------------+",
1232 ];
1233
1234 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1235 }
1236
1237 #[tokio::test]
1238 async fn test_not_equals_rule() {
1239 let df = create_test_df().await;
1240 let rule = dfq_not_eq(lit(25));
1241 let result = rule.apply(df, "age").unwrap();
1242
1243 let expected = vec![
1244 "+----+---------+-----+-------+----------------+",
1245 "| id | name | age | score | age_not_equals |",
1246 "+----+---------+-----+-------+----------------+",
1247 "| 1 | Alice | 25 | 85.5 | false |",
1248 "| 2 | Bob | | 92.0 | |",
1249 "| 3 | Charlie | 30 | | true |",
1250 "+----+---------+-----+-------+----------------+",
1251 ];
1252
1253 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1254 }
1255
1256 #[tokio::test]
1257 async fn test_string_length_rules() {
1258 let schema = Schema::new(vec![
1260 Field::new("id", DataType::Int32, false),
1261 Field::new("text", DataType::Utf8, true),
1262 ]);
1263
1264 let batch = RecordBatch::try_new(
1265 Arc::new(schema),
1266 vec![
1267 Arc::new(Int32Array::from(vec![1, 2, 3, 4, 5])),
1268 Arc::new(StringArray::from(vec![
1269 Some(""), Some("a"), Some("abc"), Some("abcdef"), None, ])),
1275 ],
1276 )
1277 .unwrap();
1278
1279 let ctx = SessionContext::new();
1280 let df = ctx.read_batch(batch).unwrap();
1281
1282 let rule = dfq_str_length(Some(2), Some(5));
1284 let result = rule.apply(df.clone(), "text").unwrap();
1285
1286 let expected = vec![
1287 "+----+--------+-------------+",
1288 "| id | text | text_length |",
1289 "+----+--------+-------------+",
1290 "| 1 | | false |",
1291 "| 2 | a | false |",
1292 "| 3 | abc | true |",
1293 "| 4 | abcdef | false |",
1294 "| 5 | | |",
1295 "+----+--------+-------------+",
1296 ];
1297
1298 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1299
1300 let rule = dfq_str_min_length(3);
1302 let result = rule.apply(df.clone(), "text").unwrap();
1303
1304 let expected = vec![
1305 "+----+--------+-------------+",
1306 "| id | text | text_length |",
1307 "+----+--------+-------------+",
1308 "| 1 | | false |",
1309 "| 2 | a | false |",
1310 "| 3 | abc | true |",
1311 "| 4 | abcdef | true |",
1312 "| 5 | | |",
1313 "+----+--------+-------------+",
1314 ];
1315
1316 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1317
1318 let rule = dfq_str_max_length(3);
1320 let result = rule.apply(df.clone(), "text").unwrap();
1321
1322 let expected = vec![
1323 "+----+--------+-------------+",
1324 "| id | text | text_length |",
1325 "+----+--------+-------------+",
1326 "| 1 | | true |",
1327 "| 2 | a | true |",
1328 "| 3 | abc | false |",
1329 "| 4 | abcdef | false |",
1330 "| 5 | | |",
1331 "+----+--------+-------------+",
1332 ];
1333
1334 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1335
1336 let rule = dfq_str_empty();
1338 let result = rule.apply(df.clone(), "text").unwrap();
1339
1340 let expected = vec![
1341 "+----+--------+-------------+",
1342 "| id | text | text_length |",
1343 "+----+--------+-------------+",
1344 "| 1 | | false |",
1345 "| 2 | a | false |",
1346 "| 3 | abc | false |",
1347 "| 4 | abcdef | false |",
1348 "| 5 | | |",
1349 "+----+--------+-------------+",
1350 ];
1351
1352 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1353
1354 let rule = dfq_str_not_empty();
1356 let result = rule.apply(df, "text").unwrap();
1357
1358 let expected = vec![
1359 "+----+--------+-------------+",
1360 "| id | text | text_length |",
1361 "+----+--------+-------------+",
1362 "| 1 | | false |",
1363 "| 2 | a | true |",
1364 "| 3 | abc | true |",
1365 "| 4 | abcdef | true |",
1366 "| 5 | | |",
1367 "+----+--------+-------------+",
1368 ];
1369
1370 assert_batches_eq!(&expected, &result.collect().await.unwrap());
1371 }
1372}