Skip to main content

datafusion_quality/rules/
column.rs

1use crate::{ColumnRule, ValidationError, error::DataFusionSnafu};
2use datafusion::{logical_expr::Between, prelude::*};
3use snafu::ResultExt;
4use std::sync::Arc;
5
6/// Rule that checks if values in a column are not null
7#[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
48/// Creates a rule that checks if values in a column are not null.
49///
50/// # Arguments
51///
52/// * `column_name` - The name of the column to check
53///
54/// # Examples
55///
56/// ```
57/// use datafusion_quality::rules::column::dfq_not_null;
58/// use datafusion_quality::RuleSet;
59///
60/// // Create a rule to check if age column is not null
61/// let rule = dfq_not_null();
62/// let mut ruleset = RuleSet::new();
63/// ruleset.with_column_rule("age", rule);
64/// ```
65pub fn dfq_not_null() -> Arc<NullRule> {
66    Arc::new(NullRule::new(Some(true)))
67}
68
69/// Creates a rule that checks if values in a column are null.
70///
71/// # Arguments
72///
73/// * `column_name` - The name of the column to check
74///
75/// # Examples
76///
77/// ```
78/// use datafusion_quality::rules::column::dfq_null;
79/// use datafusion_quality::RuleSet;
80///
81/// // Create a rule to check if age column is null
82/// let rule = dfq_null();
83/// let mut ruleset = RuleSet::new();
84/// ruleset.with_column_rule("age", rule);
85/// ```
86pub fn dfq_null() -> Arc<NullRule> {
87    Arc::new(NullRule::new(Some(false)))
88}
89
90/// Rule that checks if values in a column fall within a specified range
91#[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
135/// Creates a rule that checks if values in a column fall within a specified range.
136///
137/// # Arguments
138///
139/// * `min` - The minimum value of the range (inclusive)
140/// * `max` - The maximum value of the range (inclusive)
141///
142/// # Examples
143///
144/// ```
145/// use datafusion_quality::rules::column::dfq_in_range;
146/// use datafusion_quality::RuleSet;
147///
148/// // Create a rule to check if score is between 0 and 100
149/// let rule = dfq_in_range(0.0, 100.0);
150/// let mut ruleset = RuleSet::new();
151/// ruleset.with_column_rule("score", rule);
152/// ```
153pub fn dfq_in_range(min: f64, max: f64) -> Arc<RangeRule> {
154    Arc::new(RangeRule::new(min, max, None))
155}
156
157/// Creates a rule that checks if values in a column do not fall within a specified range.
158///
159/// # Arguments
160///
161/// * `min` - The minimum value of the range (inclusive)
162/// * `max` - The maximum value of the range (inclusive)
163///
164/// # Examples
165///
166/// ```
167/// use datafusion_quality::rules::column::dfq_not_in_range;
168/// use datafusion_quality::RuleSet;
169///
170/// // Create a rule to check if score is not between 0 and 100
171/// let rule = dfq_not_in_range(0.0, 100.0);
172/// let mut ruleset = RuleSet::new();
173/// ruleset.with_column_rule("score", rule);
174/// ```
175pub fn dfq_not_in_range(min: f64, max: f64) -> Arc<RangeRule> {
176    Arc::new(RangeRule::new(min, max, Some(true)))
177}
178
179/// Rule that checks if values in a column match a pattern
180#[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
235/// Creates a rule that checks if values in a column match a pattern (case-sensitive).
236///
237/// # Arguments
238///
239/// * `pattern` - The SQL LIKE pattern to match against
240///
241/// # Examples
242///
243/// ```
244/// use datafusion_quality::rules::column::dfq_like;
245/// use datafusion_quality::RuleSet;
246///
247/// // Create a rule to check if name starts with 'A'
248/// let rule = dfq_like("A%");
249/// let mut ruleset = RuleSet::new();
250/// ruleset.with_column_rule("name", rule);
251/// ```
252pub fn dfq_like(pattern: &str) -> Arc<PatternRule> {
253    Arc::new(PatternRule::new(pattern, Some(false), Some(true)))
254}
255
256/// Creates a rule that checks if values in a column do not match a pattern (case-sensitive).
257///
258/// # Arguments
259///
260/// * `pattern` - The SQL LIKE pattern to match against
261///
262/// # Examples
263///
264/// ```
265/// use datafusion_quality::rules::column::dfq_not_like;
266/// use datafusion_quality::RuleSet;
267///
268/// // Create a rule to check if name does not start with 'A'
269/// let rule = dfq_not_like("A%");
270/// let mut ruleset = RuleSet::new();
271/// ruleset.with_column_rule("name", rule);
272/// ```
273pub fn dfq_not_like(pattern: &str) -> Arc<PatternRule> {
274    Arc::new(PatternRule::new(pattern, Some(true), Some(true)))
275}
276
277/// Creates a rule that checks if values in a column match a pattern (case-insensitive).
278///
279/// # Arguments
280///
281/// * `pattern` - The SQL LIKE pattern to match against
282///
283/// # Examples
284///
285/// ```
286/// use datafusion_quality::rules::column::dfq_ilike;
287/// use datafusion_quality::RuleSet;
288///
289/// // Create a rule to check if name starts with 'a' (case-insensitive)
290/// let rule = dfq_ilike("a%");
291/// let mut ruleset = RuleSet::new();
292/// ruleset.with_column_rule("name", rule);
293/// ```
294pub fn dfq_ilike(pattern: &str) -> Arc<PatternRule> {
295    Arc::new(PatternRule::new(pattern, Some(false), Some(false)))
296}
297
298/// Creates a rule that checks if values in a column do not match a pattern (case-insensitive).
299///
300/// # Arguments
301///
302/// * `pattern` - The SQL LIKE pattern to match against
303///
304/// # Examples
305///
306/// ```
307/// use datafusion_quality::rules::column::dfq_not_ilike;
308/// use datafusion_quality::RuleSet;
309///
310/// // Create a rule to check if name does not start with 'a' (case-insensitive)
311/// let rule = dfq_not_ilike("a%");
312/// let mut ruleset = RuleSet::new();
313/// ruleset.with_column_rule("name", rule);
314/// ```
315pub fn dfq_not_ilike(pattern: &str) -> Arc<PatternRule> {
316    Arc::new(PatternRule::new(pattern, Some(true), Some(false)))
317}
318
319/// Rule that checks if values in a column are less than a value
320#[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
391/// Creates a rule that checks if values in a column are less than a value.
392///
393/// # Arguments
394///
395/// * `value` - The value to compare against
396///
397/// # Examples
398///
399/// ```
400/// use datafusion_quality::rules::column::dfq_lt;
401/// use datafusion_quality::RuleSet;
402/// use datafusion::prelude::*;
403///
404/// // Create a rule to check if age is less than 30
405/// let rule = dfq_lt(lit(30));
406/// let mut ruleset = RuleSet::new();
407/// ruleset.with_column_rule("age", rule);
408/// ```
409pub fn dfq_lt(value: Expr) -> Arc<ComparisonRule> {
410    Arc::new(ComparisonRule::new(
411        value,
412        false,
413        false,
414        ComparisonType::LessThan,
415    ))
416}
417
418/// Creates a rule that checks if values in a column are less than or equal to a value.
419///
420/// # Arguments
421///
422/// * `value` - The value to compare against
423///
424/// # Examples
425///
426/// ```
427/// use datafusion_quality::rules::column::dfq_lte;
428/// use datafusion_quality::RuleSet;
429/// use datafusion::prelude::*;
430///
431/// // Create a rule to check if age is less than or equal to 30
432/// let rule = dfq_lte(lit(30));
433/// let mut ruleset = RuleSet::new();
434/// ruleset.with_column_rule("age", rule);
435/// ```
436pub fn dfq_lte(value: Expr) -> Arc<ComparisonRule> {
437    Arc::new(ComparisonRule::new(
438        value,
439        false,
440        true,
441        ComparisonType::LessThan,
442    ))
443}
444
445/// Creates a rule that checks if values in a column are not less than a value.
446///
447/// # Arguments
448///
449/// * `value` - The value to compare against
450///
451/// # Examples
452///
453/// ```
454/// use datafusion_quality::rules::column::dfq_not_lt;
455/// use datafusion_quality::RuleSet;
456/// use datafusion::prelude::*;
457///
458/// // Create a rule to check if age is not less than 30
459/// let rule = dfq_not_lt(lit(30));
460/// let mut ruleset = RuleSet::new();
461/// ruleset.with_column_rule("age", rule);
462/// ```
463pub 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
472/// Creates a rule that checks if values in a column are not less than or equal to a value.
473///
474/// # Arguments
475///
476/// * `value` - The value to compare against
477///
478/// # Examples
479///
480/// ```
481/// use datafusion_quality::rules::column::dfq_not_lte;
482/// use datafusion_quality::RuleSet;
483/// use datafusion::prelude::*;
484///
485/// // Create a rule to check if age is not less than or equal to 30
486/// let rule = dfq_not_lte(lit(30));
487/// let mut ruleset = RuleSet::new();
488/// ruleset.with_column_rule("age", rule);
489/// ```
490pub 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
499/// Creates a rule that checks if values in a column are greater than a value.
500///
501/// # Arguments
502///
503/// * `value` - The value to compare against
504///
505/// # Examples
506///
507/// ```
508/// use datafusion_quality::rules::column::dfq_gt;
509/// use datafusion_quality::RuleSet;
510/// use datafusion::prelude::*;
511///
512/// // Create a rule to check if age is greater than 25
513/// let rule = dfq_gt(lit(25));
514/// let mut ruleset = RuleSet::new();
515/// ruleset.with_column_rule("age", rule);
516/// ```
517pub fn dfq_gt(value: Expr) -> Arc<ComparisonRule> {
518    Arc::new(ComparisonRule::new(
519        value,
520        false,
521        false,
522        ComparisonType::GreaterThan,
523    ))
524}
525
526/// Creates a rule that checks if values in a column are greater than or equal to a value.
527///
528/// # Arguments
529///
530/// * `value` - The value to compare against
531///
532/// # Examples
533///
534/// ```
535/// use datafusion_quality::rules::column::dfq_gte;
536/// use datafusion_quality::RuleSet;
537/// use datafusion::prelude::*;
538///
539/// // Create a rule to check if age is greater than or equal to 25
540/// let rule = dfq_gte(lit(25));
541/// let mut ruleset = RuleSet::new();
542/// ruleset.with_column_rule("age", rule);
543/// ```
544pub fn dfq_gte(value: Expr) -> Arc<ComparisonRule> {
545    Arc::new(ComparisonRule::new(
546        value,
547        false,
548        true,
549        ComparisonType::GreaterThan,
550    ))
551}
552
553/// Creates a rule that checks if values in a column are not greater than a value.
554///
555/// # Arguments
556///
557/// * `value` - The value to compare against
558///
559/// # Examples
560///
561/// ```
562/// use datafusion_quality::rules::column::dfq_not_gt;
563/// use datafusion_quality::RuleSet;
564/// use datafusion::prelude::*;
565///
566/// // Create a rule to check if age is not greater than 25
567/// let rule = dfq_not_gt(lit(25));
568/// let mut ruleset = RuleSet::new();
569/// ruleset.with_column_rule("age", rule);
570/// ```
571pub 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
580/// Creates a rule that checks if values in a column are not greater than or equal to a value.
581///
582/// # Arguments
583///
584/// * `value` - The value to compare against
585///
586/// # Examples
587///
588/// ```
589/// use datafusion_quality::rules::column::dfq_not_gte;
590/// use datafusion_quality::RuleSet;
591/// use datafusion::prelude::*;
592///
593/// // Create a rule to check if age is not greater than or equal to 25
594/// let rule = dfq_not_gte(lit(25));
595/// let mut ruleset = RuleSet::new();
596/// ruleset.with_column_rule("age", rule);
597/// ```
598pub 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
607/// Creates a rule that checks if values in a column are equal to a value.
608///
609/// # Arguments
610///
611/// * `value` - The value to compare against
612///
613/// # Examples
614///
615/// ```
616/// use datafusion_quality::rules::column::dfq_eq;
617/// use datafusion_quality::RuleSet;
618/// use datafusion::prelude::*;
619///
620/// // Create a rule to check if age is equal to 25
621/// let rule = dfq_eq(lit(25));
622/// let mut ruleset = RuleSet::new();
623/// ruleset.with_column_rule("age", rule);
624/// ```
625pub fn dfq_eq(value: Expr) -> Arc<ComparisonRule> {
626    Arc::new(ComparisonRule::new(
627        value,
628        false,
629        false,
630        ComparisonType::Equals,
631    ))
632}
633
634/// Creates a rule that checks if values in a column are not equal to a value.
635///
636/// # Arguments
637///
638/// * `value` - The value to compare against
639///
640/// # Examples
641///
642/// ```
643/// use datafusion_quality::rules::column::dfq_not_eq;
644/// use datafusion_quality::RuleSet;
645/// use datafusion::prelude::*;
646///
647/// // Create a rule to check if age is not equal to 25
648/// let rule = dfq_not_eq(lit(25));
649/// let mut ruleset = RuleSet::new();
650/// ruleset.with_column_rule("age", rule);
651/// ```
652pub 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
716/// Creates a rule that checks if the length of a string column is between a minimum and maximum value.
717///
718/// # Arguments
719///
720/// * `min` - Optional minimum length (inclusive)
721/// * `max` - Optional maximum length (inclusive)
722///
723/// # Examples
724///
725/// ```
726/// use datafusion_quality::rules::column::dfq_str_length;
727/// use datafusion_quality::RuleSet;
728///
729/// // Create a rule to check if name length is between 3 and 10 characters
730/// let rule = dfq_str_length(Some(3), Some(10));
731/// let mut ruleset = RuleSet::new();
732/// ruleset.with_column_rule("name", rule);
733/// ```
734pub fn dfq_str_length(min: Option<u32>, max: Option<u32>) -> Arc<LengthRule> {
735    Arc::new(LengthRule::new(min, max))
736}
737
738/// Creates a rule that checks if the length of a string column is at least a minimum value.
739///
740/// # Arguments
741///
742/// * `min` - Minimum length (inclusive)
743///
744/// # Examples
745///
746/// ```
747/// use datafusion_quality::rules::column::dfq_str_min_length;
748/// use datafusion_quality::RuleSet;
749///
750/// // Create a rule to check if name is at least 3 characters long
751/// let rule = dfq_str_min_length(3);
752/// let mut ruleset = RuleSet::new();
753/// ruleset.with_column_rule("name", rule);
754/// ```
755pub fn dfq_str_min_length(min: u32) -> Arc<LengthRule> {
756    Arc::new(LengthRule::new(Some(min), None))
757}
758
759/// Creates a rule that checks if the length of a string column is at most a maximum value.
760///
761/// # Arguments
762///
763/// * `max` - Maximum length (inclusive)
764///
765/// # Examples
766///
767/// ```
768/// use datafusion_quality::rules::column::dfq_str_max_length;
769/// use datafusion_quality::RuleSet;
770///
771/// // Create a rule to check if name is at most 10 characters long
772/// let rule = dfq_str_max_length(10);
773/// let mut ruleset = RuleSet::new();
774/// ruleset.with_column_rule("name", rule);
775/// ```
776pub fn dfq_str_max_length(max: u32) -> Arc<LengthRule> {
777    Arc::new(LengthRule::new(None, Some(max)))
778}
779
780/// Creates a rule that checks if a string column is empty (length = 0).
781///
782/// # Examples
783///
784/// ```
785/// use datafusion_quality::rules::column::dfq_str_empty;
786/// use datafusion_quality::RuleSet;
787///
788/// // Create a rule to check if name is empty
789/// let rule = dfq_str_empty();
790/// let mut ruleset = RuleSet::new();
791/// ruleset.with_column_rule("name", rule);
792/// ```
793pub fn dfq_str_empty() -> Arc<LengthRule> {
794    Arc::new(LengthRule::new(None, Some(0)))
795}
796
797/// Creates a rule that checks if a string column is not empty (length > 0).
798///
799/// # Examples
800///
801/// ```
802/// use datafusion_quality::rules::column::dfq_str_not_empty;
803/// use datafusion_quality::RuleSet;
804///
805/// // Create a rule to check if name is not empty
806/// let rule = dfq_str_not_empty();
807/// let mut ruleset = RuleSet::new();
808/// ruleset.with_column_rule("name", rule);
809/// ```
810pub fn dfq_str_not_empty() -> Arc<LengthRule> {
811    Arc::new(LengthRule::new(Some(1), None))
812}
813
814/// Rule that applies a custom SQL expression to a column
815#[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
850/// Creates a rule that applies a custom SQL expression to a column.
851///
852/// # Arguments
853///
854/// * `rule_name` - A name for the custom rule
855/// * `expression` - The SQL expression to apply
856///
857/// # Examples
858///
859/// ```
860/// use datafusion_quality::rules::column::dfq_custom;
861/// use datafusion_quality::RuleSet;
862/// use datafusion::prelude::*;
863///
864/// // Create a custom rule to check if age is greater than 25
865/// let rule = dfq_custom("age_gt_25", col("age").gt(lit(25)));
866/// let mut ruleset = RuleSet::new();
867/// ruleset.with_column_rule("age", rule);
868/// ```
869pub 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        // Test negated not null rule
923        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        // Test negated range rule
959        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        // Test case sensitive pattern match
979        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        // Test case insensitive pattern match
996        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        // Test negated case sensitive pattern match
1013        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        // Test negated case insensitive pattern match
1030        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        // Create a test dataframe with strings of various lengths
1259        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(""),       // empty string
1270                    Some("a"),      // length 1
1271                    Some("abc"),    // length 3
1272                    Some("abcdef"), // length 6
1273                    None,           // null
1274                ])),
1275            ],
1276        )
1277        .unwrap();
1278
1279        let ctx = SessionContext::new();
1280        let df = ctx.read_batch(batch).unwrap();
1281
1282        // Test dfq_str_length with min=2 and max=5
1283        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        // Test dfq_str_min_length with min=3
1301        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        // Test dfq_str_max_length with max=3
1319        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        // Test dfq_str_empty
1337        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        // Test dfq_str_not_empty
1355        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}