Skip to main content

proof_of_sql_planner/
expr.rs

1use super::{
2    column_to_column_ref, placeholder_to_placeholder_expr, scalar_value_to_literal_value,
3    PlannerError, PlannerResult,
4};
5use datafusion::logical_expr::{
6    expr::{Alias, Between, Cast, InList, Placeholder},
7    BinaryExpr, Expr, Operator,
8};
9use indexmap::IndexSet;
10use proof_of_sql::{
11    base::database::{ColumnType, LiteralValue},
12    sql::{
13        proof_exprs::{DynProofExpr, ProofExpr},
14        scale_cast_binary_op,
15    },
16};
17use sqlparser::ast::Ident;
18
19/// Recursively extract all column identifiers referenced in an expression
20pub(crate) fn get_column_idents_from_expr(expr: &Expr) -> IndexSet<Ident> {
21    match expr {
22        Expr::Column(col) => {
23            let mut set = IndexSet::new();
24            set.insert(col.name.as_str().into());
25            set
26        }
27        Expr::BinaryExpr(BinaryExpr { left, right, .. }) => {
28            let mut left_idents = get_column_idents_from_expr(left);
29            left_idents.extend(get_column_idents_from_expr(right));
30            left_idents
31        }
32        Expr::Not(inner) => get_column_idents_from_expr(inner),
33        Expr::InList(InList { expr, list, .. }) => {
34            let mut idents = get_column_idents_from_expr(expr);
35            for value in list {
36                idents.extend(get_column_idents_from_expr(value));
37            }
38            idents
39        }
40        Expr::Alias(Alias { expr, .. }) | Expr::Cast(Cast { expr, .. }) => {
41            get_column_idents_from_expr(expr)
42        }
43        Expr::AggregateFunction(agg) => agg
44            .args
45            .iter()
46            .flat_map(get_column_idents_from_expr)
47            .collect(),
48        Expr::Between(Between {
49            expr, low, high, ..
50        }) => {
51            let mut idents = get_column_idents_from_expr(expr);
52            idents.extend(get_column_idents_from_expr(low));
53            idents.extend(get_column_idents_from_expr(high));
54            idents
55        }
56        _ => IndexSet::new(),
57    }
58}
59
60/// Convert a [`BinaryExpr`] to [`DynProofExpr`]
61fn binary_expr_to_proof_expr(
62    left: &Expr,
63    right: &Expr,
64    op: Operator,
65    schema: &[(Ident, ColumnType)],
66) -> PlannerResult<DynProofExpr> {
67    let left_proof_expr = expr_to_proof_expr(left, schema)?;
68    let right_proof_expr = expr_to_proof_expr(right, schema)?;
69    binary_proof_exprs_to_proof_expr(left_proof_expr, right_proof_expr, op)
70}
71
72/// Apply a binary [`Operator`] to two already-converted [`DynProofExpr`]s.
73///
74/// Scale-casting is performed here for operators that require it, so callers
75/// that already hold `DynProofExpr`s can skip the `Expr` conversion step.
76#[expect(
77    clippy::missing_panics_doc,
78    reason = "Output of comparisons is always boolean"
79)]
80fn binary_proof_exprs_to_proof_expr(
81    left_proof_expr: DynProofExpr,
82    right_proof_expr: DynProofExpr,
83    op: Operator,
84) -> PlannerResult<DynProofExpr> {
85    let (left_proof_expr, right_proof_expr) = match op {
86        Operator::Eq
87        | Operator::NotEq
88        | Operator::Lt
89        | Operator::Gt
90        | Operator::LtEq
91        | Operator::GtEq
92        | Operator::Plus
93        | Operator::Minus => scale_cast_binary_op(left_proof_expr, right_proof_expr)?,
94        _ => (left_proof_expr, right_proof_expr),
95    };
96
97    match op {
98        Operator::And => Ok(DynProofExpr::try_new_and(
99            left_proof_expr,
100            right_proof_expr,
101        )?),
102        Operator::Or => Ok(DynProofExpr::try_new_or(left_proof_expr, right_proof_expr)?),
103        Operator::Multiply => Ok(DynProofExpr::try_new_multiply(
104            left_proof_expr,
105            right_proof_expr,
106        )?),
107        Operator::Eq => Ok(DynProofExpr::try_new_equals(
108            left_proof_expr,
109            right_proof_expr,
110        )?),
111        Operator::NotEq => Ok(DynProofExpr::try_new_not(DynProofExpr::try_new_equals(
112            left_proof_expr,
113            right_proof_expr,
114        )?)
115        .expect("An equality expression must have a boolean data type...")),
116        Operator::Lt => Ok(DynProofExpr::try_new_inequality(
117            left_proof_expr,
118            right_proof_expr,
119            true,
120        )?),
121        Operator::Gt => Ok(DynProofExpr::try_new_inequality(
122            left_proof_expr,
123            right_proof_expr,
124            false,
125        )?),
126        Operator::LtEq => Ok(DynProofExpr::try_new_not(DynProofExpr::try_new_inequality(
127            left_proof_expr,
128            right_proof_expr,
129            false,
130        )?)
131        .expect("An inequality expression must have a boolean data type...")),
132        Operator::GtEq => Ok(DynProofExpr::try_new_not(DynProofExpr::try_new_inequality(
133            left_proof_expr,
134            right_proof_expr,
135            true,
136        )?)
137        .expect("An inequality expression must have a boolean data type...")),
138        Operator::Plus => Ok(DynProofExpr::try_new_add(
139            left_proof_expr,
140            right_proof_expr,
141        )?),
142        Operator::Minus => Ok(DynProofExpr::try_new_subtract(
143            left_proof_expr,
144            right_proof_expr,
145        )?),
146        _ => Err(PlannerError::UnsupportedBinaryOperator { op }),
147    }
148}
149
150/// Convert a `DataFusion` [`Expr`] into a provable [`DynProofExpr`], using `schema`
151/// to resolve column references and their types.
152///
153/// # Errors
154/// Returns a [`PlannerError`] if the expression (or any subexpression) is unsupported,
155/// references an unknown column, or cannot be lowered to a `DynProofExpr`.
156pub fn expr_to_proof_expr(
157    expr: &Expr,
158    schema: &[(Ident, ColumnType)],
159) -> PlannerResult<DynProofExpr> {
160    match expr {
161        Expr::Alias(Alias { expr, .. }) => expr_to_proof_expr(expr, schema),
162        Expr::Column(col) => Ok(DynProofExpr::new_column(column_to_column_ref(col, schema)?)),
163        Expr::Placeholder(placeholder) => placeholder_to_placeholder_expr(placeholder),
164        Expr::BinaryExpr(BinaryExpr { left, right, op }) => {
165            binary_expr_to_proof_expr(left, right, *op, schema)
166        }
167        Expr::Literal(val) => Ok(DynProofExpr::new_literal(scalar_value_to_literal_value(
168            val.clone(),
169        )?)),
170        Expr::Not(expr) => {
171            let proof_expr = expr_to_proof_expr(expr, schema)?;
172            Ok(DynProofExpr::try_new_not(proof_expr)?)
173        }
174        Expr::InList(InList {
175            expr,
176            list,
177            negated,
178        }) => {
179            // Lower `a IN (v_1, ..., v_k)`. The list length is uncapped in both branches.
180            //
181            //   - numeric `a`: use the product form `(a - v_1) * ... * (a - v_k) = 0`.
182            //     A field is an integral domain, so the product is zero iff some factor is.
183            //   - non-numeric `a` (e.g. varchar): the typed `-` operator rejects these, so
184            //     fall back to `a = v_1 OR ... OR a = v_k`. Equality subtracts the embedded
185            //     hashes internally, which is exactly what a membership zero-test needs.
186            let needle = expr_to_proof_expr(expr, schema)?;
187            // Split off the first value. `None` means `a IN ()`, which is always false;
188            // keeping that here (not an early return) lets the `negated` wrap below
189            // still apply, so `a NOT IN ()` correctly negates to `true`.
190            let comparison = match list.split_first() {
191                None => DynProofExpr::new_literal(LiteralValue::Boolean(false)),
192                Some((first, rest)) => {
193                    if needle.data_type().is_numeric() {
194                        // product form: (a - v_1) * ... * (a - v_k) = 0. `scale_cast_binary_op`
195                        // aligns the needle to each `v_i`'s scale before subtracting; this is a
196                        // numeric-only concern, so it stays out of the varchar branch.
197                        let term = |value: &Expr| -> PlannerResult<_> {
198                            let (n, v) = scale_cast_binary_op(
199                                needle.clone(),
200                                expr_to_proof_expr(value, schema)?,
201                            )?;
202                            Ok(DynProofExpr::try_new_subtract(n, v)?)
203                        };
204                        let product = rest.iter().try_fold(
205                            term(first)?,
206                            |acc, value| -> PlannerResult<_> {
207                                Ok(DynProofExpr::try_new_multiply(acc, term(value)?)?)
208                            },
209                        )?;
210                        let zero = DynProofExpr::new_literal(LiteralValue::BigInt(0));
211                        let (product, zero) = scale_cast_binary_op(product, zero)?;
212                        DynProofExpr::try_new_equals(product, zero)?
213                    } else {
214                        // a = v_1 OR ... OR a = v_k. The needle is non-numeric, so there is no
215                        // scale to align; `=` compares the lowered operands directly.
216                        let eq = |value: &Expr| -> PlannerResult<_> {
217                            Ok(DynProofExpr::try_new_equals(
218                                needle.clone(),
219                                expr_to_proof_expr(value, schema)?,
220                            )?)
221                        };
222                        rest.iter()
223                            .try_fold(eq(first)?, |acc, value| -> PlannerResult<_> {
224                                Ok(DynProofExpr::try_new_or(acc, eq(value)?)?)
225                            })?
226                    }
227                }
228            };
229            if *negated {
230                Ok(DynProofExpr::try_new_not(comparison)?)
231            } else {
232                Ok(comparison)
233            }
234        }
235        Expr::Cast(cast) => {
236            match &*cast.expr {
237                // handle cases such as `$1::int`
238                Expr::Placeholder(placeholder) if placeholder.data_type.is_none() => {
239                    let typed_placeholder =
240                        Placeholder::new(placeholder.id.clone(), Some(cast.data_type.clone()));
241                    placeholder_to_placeholder_expr(&typed_placeholder)
242                }
243                _ => {
244                    let from_expr = expr_to_proof_expr(&cast.expr, schema)?;
245                    let to_type = cast.data_type.clone().try_into().map_err(|_| {
246                        PlannerError::UnsupportedDataType {
247                            data_type: cast.data_type.clone(),
248                        }
249                    })?;
250                    Ok(
251                        DynProofExpr::try_new_cast(from_expr.clone(), to_type).map_or_else(
252                            |_| DynProofExpr::try_new_scaling_cast(from_expr, to_type),
253                            Ok,
254                        )?,
255                    )
256                }
257            }
258        }
259        Expr::Between(Between {
260            expr,
261            negated,
262            low,
263            high,
264        }) => between_to_proof_expr(expr, *negated, low, high, schema),
265        _ => Err(PlannerError::UnsupportedLogicalExpression {
266            expr: Box::new(expr.clone()),
267        }),
268    }
269}
270
271/// Convert a [`Between`] expression to [`DynProofExpr`]
272fn between_to_proof_expr(
273    expr: &Expr,
274    negated: bool,
275    low: &Expr,
276    high: &Expr,
277    schema: &[(Ident, ColumnType)],
278) -> PlannerResult<DynProofExpr> {
279    let expr_proof = expr_to_proof_expr(expr, schema)?;
280    let low_proof = expr_to_proof_expr(low, schema)?;
281    let high_proof = expr_to_proof_expr(high, schema)?;
282    // expr < low OR expr > high
283    let out_of_range = binary_proof_exprs_to_proof_expr(
284        binary_proof_exprs_to_proof_expr(expr_proof.clone(), low_proof, Operator::Lt)?,
285        binary_proof_exprs_to_proof_expr(expr_proof, high_proof, Operator::Gt)?,
286        Operator::Or,
287    )?;
288    if negated {
289        Ok(out_of_range)
290    } else {
291        Ok(DynProofExpr::try_new_not(out_of_range)?)
292    }
293}
294
295#[cfg(test)]
296mod tests {
297    use super::*;
298    use crate::df_util::*;
299    use arrow::datatypes::DataType;
300    use core::ops::{Add, Mul, Sub};
301    use datafusion::{
302        catalog::TableReference,
303        common::{Column, ScalarValue},
304        logical_expr::{expr::Placeholder, lit, Cast},
305    };
306    use proof_of_sql::base::{
307        database::{ColumnRef, ColumnType, LiteralValue, TableRef},
308        math::decimal::Precision,
309    };
310
311    #[expect(non_snake_case)]
312    fn COLUMN_INT() -> DynProofExpr {
313        DynProofExpr::new_column(ColumnRef::new(
314            TableRef::from_names(Some("namespace"), "table_name"),
315            "column".into(),
316            ColumnType::Int,
317        ))
318    }
319
320    #[expect(non_snake_case)]
321    fn COLUMN1_SMALLINT() -> DynProofExpr {
322        DynProofExpr::new_column(ColumnRef::new(
323            TableRef::from_names(Some("namespace"), "table_name"),
324            "column1".into(),
325            ColumnType::SmallInt,
326        ))
327    }
328
329    #[expect(non_snake_case)]
330    fn COLUMN2_BIGINT() -> DynProofExpr {
331        DynProofExpr::new_column(ColumnRef::new(
332            TableRef::from_names(Some("namespace"), "table_name"),
333            "column2".into(),
334            ColumnType::BigInt,
335        ))
336    }
337
338    #[expect(non_snake_case)]
339    fn COLUMN1_BOOLEAN() -> DynProofExpr {
340        DynProofExpr::new_column(ColumnRef::new(
341            TableRef::from_names(Some("namespace"), "table_name"),
342            "column1".into(),
343            ColumnType::Boolean,
344        ))
345    }
346
347    #[expect(non_snake_case)]
348    fn COLUMN2_BOOLEAN() -> DynProofExpr {
349        DynProofExpr::new_column(ColumnRef::new(
350            TableRef::from_names(Some("namespace"), "table_name"),
351            "column2".into(),
352            ColumnType::Boolean,
353        ))
354    }
355
356    #[expect(non_snake_case)]
357    fn COLUMN3_DECIMAL_75_5() -> DynProofExpr {
358        DynProofExpr::new_column(ColumnRef::new(
359            TableRef::from_names(Some("namespace"), "table_name"),
360            "column3".into(),
361            ColumnType::Decimal75(
362                Precision::new(75).expect("Precision is definitely valid"),
363                5,
364            ),
365        ))
366    }
367
368    #[expect(non_snake_case)]
369    fn COLUMN2_DECIMAL_25_5() -> DynProofExpr {
370        DynProofExpr::new_column(ColumnRef::new(
371            TableRef::from_names(Some("namespace"), "table_name"),
372            "column2".into(),
373            ColumnType::Decimal75(
374                Precision::new(25).expect("Precision is definitely valid"),
375                5,
376            ),
377        ))
378    }
379
380    // Alias
381    #[test]
382    fn we_can_convert_alias_to_proof_expr() {
383        // Column
384        let expr = df_column("namespace.table_name", "column").alias("alias");
385        let schema = vec![("column".into(), ColumnType::Int)];
386        assert_eq!(expr_to_proof_expr(&expr, &schema).unwrap(), COLUMN_INT());
387    }
388
389    // Column
390    #[test]
391    fn we_can_convert_column_expr_to_proof_expr() {
392        // Column
393        let expr = df_column("namespace.table_name", "column");
394        let schema = vec![("column".into(), ColumnType::Int)];
395        assert_eq!(expr_to_proof_expr(&expr, &schema).unwrap(), COLUMN_INT());
396    }
397
398    // IN
399    #[test]
400    fn we_can_convert_in_list_to_proof_expr() {
401        // `a IN (...)` lowers to `(product of differences) = 0`, i.e. a top-level equals.
402        let expr = df_column("namespace.table_name", "column")
403            .in_list(vec![lit(1_i64), lit(2_i64), lit(3_i64)], false);
404        let schema = vec![("column".into(), ColumnType::BigInt)];
405        assert!(matches!(
406            expr_to_proof_expr(&expr, &schema).unwrap(),
407            DynProofExpr::Equals(_)
408        ));
409    }
410
411    #[test]
412    fn we_can_convert_not_in_list_to_proof_expr() {
413        // `NOT IN` wraps the lowered `IN` expression in a `NOT`.
414        let expr = df_column("namespace.table_name", "column").in_list(vec![lit(1_i64)], true);
415        let schema = vec![("column".into(), ColumnType::BigInt)];
416        assert!(matches!(
417            expr_to_proof_expr(&expr, &schema).unwrap(),
418            DynProofExpr::Not(_)
419        ));
420    }
421
422    #[test]
423    fn we_convert_an_empty_in_list_to_a_false_literal() {
424        // `a IN ()` is always false.
425        let expr = df_column("namespace.table_name", "column").in_list(vec![], false);
426        let schema = vec![("column".into(), ColumnType::BigInt)];
427        assert_eq!(
428            expr_to_proof_expr(&expr, &schema).unwrap(),
429            DynProofExpr::new_literal(LiteralValue::Boolean(false))
430        );
431    }
432
433    #[test]
434    fn we_convert_an_empty_not_in_list_to_a_true_literal() {
435        // `a NOT IN ()` is always true: the empty-list `false` must still flow
436        // through the `negated` NOT wrap rather than short-circuiting.
437        let expr = df_column("namespace.table_name", "column").in_list(vec![], true);
438        let schema = vec![("column".into(), ColumnType::BigInt)];
439        assert_eq!(
440            expr_to_proof_expr(&expr, &schema).unwrap(),
441            DynProofExpr::try_new_not(DynProofExpr::new_literal(LiteralValue::Boolean(false)))
442                .unwrap()
443        );
444    }
445
446    #[test]
447    fn we_convert_a_varchar_in_list_to_an_or_chain() {
448        // `-` rejects varchar, so non-numeric `IN` falls back to `a = v_1 OR a = v_2`.
449        let expr =
450            df_column("namespace.table_name", "column").in_list(vec![lit("a"), lit("b")], false);
451        let schema = vec![("column".into(), ColumnType::VarChar)];
452        assert!(matches!(
453            expr_to_proof_expr(&expr, &schema).unwrap(),
454            DynProofExpr::Or(_)
455        ));
456    }
457
458    // BinaryExpr
459    #[test]
460    fn we_can_convert_comparison_binary_expr_to_proof_expr() {
461        let schema = vec![
462            ("column1".into(), ColumnType::SmallInt),
463            ("column2".into(), ColumnType::BigInt),
464        ];
465
466        // Eq
467        let expr = df_column("namespace.table_name", "column1")
468            .eq(df_column("namespace.table_name", "column2"));
469        assert_eq!(
470            expr_to_proof_expr(&expr, &schema).unwrap(),
471            DynProofExpr::try_new_equals(COLUMN1_SMALLINT(), COLUMN2_BIGINT()).unwrap()
472        );
473
474        // Lt
475        let expr = df_column("namespace.table_name", "column1")
476            .lt(df_column("namespace.table_name", "column2"));
477        assert_eq!(
478            expr_to_proof_expr(&expr, &schema).unwrap(),
479            DynProofExpr::try_new_inequality(COLUMN1_SMALLINT(), COLUMN2_BIGINT(), true).unwrap()
480        );
481
482        // Gt
483        let expr = df_column("namespace.table_name", "column1")
484            .gt(df_column("namespace.table_name", "column2"));
485        assert_eq!(
486            expr_to_proof_expr(&expr, &schema).unwrap(),
487            DynProofExpr::try_new_inequality(COLUMN1_SMALLINT(), COLUMN2_BIGINT(), false).unwrap()
488        );
489
490        // LtEq
491        let expr = df_column("namespace.table_name", "column1")
492            .lt_eq(df_column("namespace.table_name", "column2"));
493        assert_eq!(
494            expr_to_proof_expr(&expr, &schema).unwrap(),
495            DynProofExpr::try_new_not(
496                DynProofExpr::try_new_inequality(COLUMN1_SMALLINT(), COLUMN2_BIGINT(), false)
497                    .unwrap()
498            )
499            .unwrap()
500        );
501
502        // GtEq
503        let expr = df_column("namespace.table_name", "column1")
504            .gt_eq(df_column("namespace.table_name", "column2"));
505        assert_eq!(
506            expr_to_proof_expr(&expr, &schema).unwrap(),
507            DynProofExpr::try_new_not(
508                DynProofExpr::try_new_inequality(COLUMN1_SMALLINT(), COLUMN2_BIGINT(), true)
509                    .unwrap()
510            )
511            .unwrap()
512        );
513    }
514
515    #[expect(clippy::too_many_lines)]
516    #[test]
517    fn we_can_convert_comparison_binary_expr_to_proof_expr_with_scale_cast() {
518        let schema = vec![
519            ("column1".into(), ColumnType::SmallInt),
520            (
521                "column2".into(),
522                ColumnType::Decimal75(Precision::new(25).unwrap(), 5),
523            ),
524            (
525                "column3".into(),
526                ColumnType::Decimal75(Precision::new(75).unwrap(), 5),
527            ),
528        ];
529
530        // Eq
531        let expr = df_column("namespace.table_name", "column1")
532            .eq(df_column("namespace.table_name", "column3"));
533        assert_eq!(
534            expr_to_proof_expr(&expr, &schema).unwrap(),
535            DynProofExpr::try_new_equals(
536                DynProofExpr::try_new_scaling_cast(
537                    COLUMN1_SMALLINT(),
538                    ColumnType::Decimal75(
539                        Precision::new(10).expect("Precision is definitely valid"),
540                        5
541                    )
542                )
543                .unwrap(),
544                COLUMN3_DECIMAL_75_5()
545            )
546            .unwrap()
547        );
548
549        // Lt
550        let expr = df_column("namespace.table_name", "column1")
551            .lt(df_column("namespace.table_name", "column2"));
552        assert_eq!(
553            expr_to_proof_expr(&expr, &schema).unwrap(),
554            DynProofExpr::try_new_inequality(
555                DynProofExpr::try_new_scaling_cast(
556                    COLUMN1_SMALLINT(),
557                    ColumnType::Decimal75(
558                        Precision::new(10).expect("Precision is definitely valid"),
559                        5
560                    )
561                )
562                .unwrap(),
563                COLUMN2_DECIMAL_25_5(),
564                true
565            )
566            .unwrap()
567        );
568
569        // Gt
570        let expr = df_column("namespace.table_name", "column1")
571            .gt(df_column("namespace.table_name", "column2"));
572        assert_eq!(
573            expr_to_proof_expr(&expr, &schema).unwrap(),
574            DynProofExpr::try_new_inequality(
575                DynProofExpr::try_new_scaling_cast(
576                    COLUMN1_SMALLINT(),
577                    ColumnType::Decimal75(
578                        Precision::new(10).expect("Precision is definitely valid"),
579                        5
580                    )
581                )
582                .unwrap(),
583                COLUMN2_DECIMAL_25_5(),
584                false
585            )
586            .unwrap()
587        );
588
589        // LtEq
590        let expr = df_column("namespace.table_name", "column1")
591            .lt_eq(df_column("namespace.table_name", "column2"));
592        assert_eq!(
593            expr_to_proof_expr(&expr, &schema).unwrap(),
594            DynProofExpr::try_new_not(
595                DynProofExpr::try_new_inequality(
596                    DynProofExpr::try_new_scaling_cast(
597                        COLUMN1_SMALLINT(),
598                        ColumnType::Decimal75(
599                            Precision::new(10).expect("Precision is definitely valid"),
600                            5
601                        )
602                    )
603                    .unwrap(),
604                    COLUMN2_DECIMAL_25_5(),
605                    false
606                )
607                .unwrap()
608            )
609            .unwrap()
610        );
611
612        // GtEq
613        let expr = df_column("namespace.table_name", "column1")
614            .gt_eq(df_column("namespace.table_name", "column2"));
615        assert_eq!(
616            expr_to_proof_expr(&expr, &schema).unwrap(),
617            DynProofExpr::try_new_not(
618                DynProofExpr::try_new_inequality(
619                    DynProofExpr::try_new_scaling_cast(
620                        COLUMN1_SMALLINT(),
621                        ColumnType::Decimal75(
622                            Precision::new(10).expect("Precision is definitely valid"),
623                            5
624                        )
625                    )
626                    .unwrap(),
627                    COLUMN2_DECIMAL_25_5(),
628                    true
629                )
630                .unwrap()
631            )
632            .unwrap()
633        );
634    }
635
636    #[test]
637    fn we_can_convert_arithmetic_binary_expr_to_proof_expr() {
638        let schema = vec![
639            ("column1".into(), ColumnType::SmallInt),
640            ("column2".into(), ColumnType::BigInt),
641        ];
642
643        // Plus
644        let expr = Expr::BinaryExpr(BinaryExpr {
645            left: Box::new(df_column("namespace.table_name", "column1")),
646            right: Box::new(df_column("namespace.table_name", "column2")),
647            op: Operator::Plus,
648        });
649        assert_eq!(
650            expr_to_proof_expr(&expr, &schema).unwrap(),
651            DynProofExpr::try_new_add(COLUMN1_SMALLINT(), COLUMN2_BIGINT(),).unwrap()
652        );
653
654        // Minus
655        let expr = Expr::BinaryExpr(BinaryExpr {
656            left: Box::new(df_column("namespace.table_name", "column1")),
657            right: Box::new(df_column("namespace.table_name", "column2")),
658            op: Operator::Minus,
659        });
660        assert_eq!(
661            expr_to_proof_expr(&expr, &schema).unwrap(),
662            DynProofExpr::try_new_subtract(COLUMN1_SMALLINT(), COLUMN2_BIGINT(),).unwrap()
663        );
664
665        // Multiply
666        let expr = Expr::BinaryExpr(BinaryExpr {
667            left: Box::new(df_column("namespace.table_name", "column1")),
668            right: Box::new(df_column("namespace.table_name", "column2")),
669            op: Operator::Multiply,
670        });
671        assert_eq!(
672            expr_to_proof_expr(&expr, &schema).unwrap(),
673            DynProofExpr::try_new_multiply(COLUMN1_SMALLINT(), COLUMN2_BIGINT(),).unwrap()
674        );
675    }
676
677    #[test]
678    fn we_can_convert_arithmetic_binary_expr_to_proof_expr_with_scale_cast() {
679        let schema = vec![
680            ("column1".into(), ColumnType::SmallInt),
681            (
682                "column2".into(),
683                ColumnType::Decimal75(Precision::new(25).unwrap(), 5),
684            ),
685            (
686                "column3".into(),
687                ColumnType::Decimal75(Precision::new(75).unwrap(), 5),
688            ),
689        ];
690
691        // Add
692        let expr = df_column("namespace.table_name", "column1")
693            .add(df_column("namespace.table_name", "column2"));
694        assert_eq!(
695            expr_to_proof_expr(&expr, &schema).unwrap(),
696            DynProofExpr::try_new_add(
697                DynProofExpr::try_new_scaling_cast(
698                    COLUMN1_SMALLINT(),
699                    ColumnType::Decimal75(
700                        Precision::new(10).expect("Precision is definitely valid"),
701                        5
702                    )
703                )
704                .unwrap(),
705                COLUMN2_DECIMAL_25_5()
706            )
707            .unwrap()
708        );
709
710        // Subtract
711        let expr = df_column("namespace.table_name", "column1")
712            .sub(df_column("namespace.table_name", "column2"));
713        assert_eq!(
714            expr_to_proof_expr(&expr, &schema).unwrap(),
715            DynProofExpr::try_new_subtract(
716                DynProofExpr::try_new_scaling_cast(
717                    COLUMN1_SMALLINT(),
718                    ColumnType::Decimal75(
719                        Precision::new(10).expect("Precision is definitely valid"),
720                        5
721                    )
722                )
723                .unwrap(),
724                COLUMN2_DECIMAL_25_5()
725            )
726            .unwrap()
727        );
728
729        // Multiply - No scale cast!
730        let expr = df_column("namespace.table_name", "column1")
731            .mul(df_column("namespace.table_name", "column2"));
732        assert_eq!(
733            expr_to_proof_expr(&expr, &schema).unwrap(),
734            DynProofExpr::try_new_multiply(COLUMN1_SMALLINT(), COLUMN2_DECIMAL_25_5()).unwrap()
735        );
736    }
737
738    #[test]
739    fn we_can_convert_logical_binary_expr_to_proof_expr() {
740        let schema = vec![
741            ("column1".into(), ColumnType::Boolean),
742            ("column2".into(), ColumnType::Boolean),
743        ];
744
745        // And
746        let expr = df_column("namespace.table_name", "column1")
747            .and(df_column("namespace.table_name", "column2"));
748        assert_eq!(
749            expr_to_proof_expr(&expr, &schema).unwrap(),
750            DynProofExpr::try_new_and(COLUMN1_BOOLEAN(), COLUMN2_BOOLEAN()).unwrap()
751        );
752
753        // Or
754        let expr = df_column("namespace.table_name", "column1")
755            .or(df_column("namespace.table_name", "column2"));
756        assert_eq!(
757            expr_to_proof_expr(&expr, &schema).unwrap(),
758            DynProofExpr::try_new_or(COLUMN1_BOOLEAN(), COLUMN2_BOOLEAN()).unwrap()
759        );
760    }
761
762    #[test]
763    fn we_can_convert_logical_not_eq_to_proof_expr() {
764        let schema = vec![
765            ("column1".into(), ColumnType::BigInt),
766            ("column2".into(), ColumnType::BigInt),
767        ];
768
769        let expr = df_column("namespace.table_name", "column1")
770            .not_eq(df_column("namespace.table_name", "column2"));
771        assert_eq!(
772            expr_to_proof_expr(&expr, &schema).unwrap(),
773            DynProofExpr::try_new_not(
774                DynProofExpr::try_new_equals(
775                    DynProofExpr::new_column(ColumnRef::new(
776                        TableRef::from_names(Some("namespace"), "table_name"),
777                        "column1".into(),
778                        ColumnType::BigInt,
779                    )),
780                    DynProofExpr::new_column(ColumnRef::new(
781                        TableRef::from_names(Some("namespace"), "table_name"),
782                        "column2".into(),
783                        ColumnType::BigInt,
784                    ))
785                )
786                .unwrap()
787            )
788            .unwrap()
789        );
790    }
791
792    #[test]
793    fn we_cannot_convert_unsupported_binary_expr_to_proof_expr() {
794        // Unsupported binary operator
795        let expr = Expr::BinaryExpr(BinaryExpr {
796            left: Box::new(df_column("namespace.table_name", "column1")),
797            right: Box::new(df_column("namespace.table_name", "column2")),
798            op: Operator::AtArrow,
799        });
800        let schema = vec![
801            ("column1".into(), ColumnType::Boolean),
802            ("column2".into(), ColumnType::Boolean),
803        ];
804        assert!(matches!(
805            expr_to_proof_expr(&expr, &schema),
806            Err(PlannerError::UnsupportedBinaryOperator { .. })
807        ));
808    }
809
810    // Literal
811    #[test]
812    fn we_can_convert_literal_expr_to_proof_expr() {
813        let expr = Expr::Literal(ScalarValue::Int32(Some(1)));
814        assert_eq!(
815            expr_to_proof_expr(&expr, &Vec::new()).unwrap(),
816            DynProofExpr::new_literal(LiteralValue::Int(1))
817        );
818    }
819
820    // Not
821    #[test]
822    fn we_can_convert_not_expr_to_proof_expr() {
823        let expr = Expr::Not(Box::new(df_column("table_name", "column")));
824        let schema = vec![("column".into(), ColumnType::Boolean)];
825        assert_eq!(
826            expr_to_proof_expr(&expr, &schema).unwrap(),
827            DynProofExpr::try_new_not(DynProofExpr::new_column(ColumnRef::new(
828                TableRef::from_names(None, "table_name"),
829                "column".into(),
830                ColumnType::Boolean
831            )))
832            .unwrap()
833        );
834    }
835
836    // Cast
837    #[test]
838    fn we_can_convert_cast_expr_to_proof_expr() {
839        let expr = Expr::Cast(Cast::new(
840            Box::new(Expr::Literal(ScalarValue::Boolean(Some(true)))),
841            DataType::Int32,
842        ));
843        let expression = expr_to_proof_expr(&expr, &Vec::new()).unwrap();
844        assert_eq!(
845            expression,
846            DynProofExpr::try_new_cast(
847                DynProofExpr::new_literal(LiteralValue::Boolean(true)),
848                ColumnType::Int
849            )
850            .unwrap()
851        );
852    }
853
854    #[test]
855    fn we_cannot_convert_cast_expr_to_proof_expr_when_inner_expr_to_proof_expr_fails() {
856        // Unsupported logical expression
857        let expr = Expr::Cast(Cast::new(
858            Box::new(Expr::Literal(ScalarValue::UInt64(Some(100)))),
859            DataType::Int16,
860        ));
861        let expression = expr_to_proof_expr(&expr, &Vec::new()).unwrap_err();
862        assert!(matches!(
863            expression,
864            PlannerError::UnsupportedDataType { data_type: _ }
865        ));
866    }
867
868    #[test]
869    fn we_cannot_convert_cast_expr_to_proof_expr_for_unsupported_datatypes() {
870        // Unsupported logical expression
871        let expr = Expr::Cast(Cast::new(
872            Box::new(Expr::Literal(ScalarValue::Boolean(Some(true)))),
873            DataType::UInt16,
874        ));
875        let expression = expr_to_proof_expr(&expr, &Vec::new()).unwrap_err();
876        assert!(matches!(
877            expression,
878            PlannerError::UnsupportedDataType { data_type: _ }
879        ));
880    }
881
882    #[test]
883    fn we_cannot_convert_cast_expr_to_proof_expr_for_datatypes_for_which_casting_is_not_supported()
884    {
885        // Unsupported logical expression
886        let expr = Expr::Cast(Cast::new(
887            Box::new(Expr::Literal(ScalarValue::Int16(Some(100)))),
888            DataType::Boolean,
889        ));
890        let expression = expr_to_proof_expr(&expr, &Vec::new()).unwrap_err();
891        assert!(matches!(
892            expression,
893            PlannerError::AnalyzeError { source: _ }
894        ));
895    }
896
897    // Placeholder
898    #[test]
899    fn we_can_convert_placeholder_to_proof_expr() {
900        let expr = Expr::Placeholder(Placeholder {
901            id: "$1".to_string(),
902            data_type: Some(DataType::Int32),
903        });
904        let expression = expr_to_proof_expr(&expr, &Vec::new()).unwrap();
905        assert_eq!(
906            expression,
907            DynProofExpr::try_new_placeholder(1, ColumnType::Int).unwrap()
908        );
909    }
910
911    // Placeholder with data type specified by cast
912    #[test]
913    fn we_can_convert_placeholder_with_data_type_specified_by_cast_to_proof_expr() {
914        let expr = Expr::Cast(Cast::new(
915            Box::new(Expr::Placeholder(Placeholder {
916                id: "$1".to_string(),
917                data_type: None,
918            })),
919            DataType::Int32,
920        ));
921        let expression = expr_to_proof_expr(&expr, &Vec::new()).unwrap();
922        assert_eq!(
923            expression,
924            DynProofExpr::try_new_placeholder(1, ColumnType::Int).unwrap()
925        );
926    }
927
928    // Unsupported logical expression
929    #[test]
930    fn we_cannot_convert_unsupported_expr_to_proof_expr() {
931        let expr = Expr::OuterReferenceColumn(
932            DataType::Int32,
933            Column::new(None::<TableReference>, "column"),
934        );
935        assert!(matches!(
936            expr_to_proof_expr(&expr, &Vec::new()),
937            Err(PlannerError::UnsupportedLogicalExpression { .. })
938        ));
939    }
940
941    // Between
942    #[test]
943    fn we_can_convert_between_expr_to_proof_expr() {
944        let schema = vec![("column1".into(), ColumnType::BigInt)];
945
946        let col = df_column("namespace.table_name", "column1");
947        let low = Expr::Literal(ScalarValue::Int64(Some(10)));
948        let high = Expr::Literal(ScalarValue::Int64(Some(20)));
949        let expr = col.between(low, high);
950
951        let col_expr = DynProofExpr::new_column(ColumnRef::new(
952            TableRef::from_names(Some("namespace"), "table_name"),
953            "column1".into(),
954            ColumnType::BigInt,
955        ));
956        let low_expr = DynProofExpr::new_literal(LiteralValue::BigInt(10));
957        let high_expr = DynProofExpr::new_literal(LiteralValue::BigInt(20));
958
959        let expected = DynProofExpr::try_new_not(
960            DynProofExpr::try_new_or(
961                DynProofExpr::try_new_inequality(col_expr.clone(), low_expr, true).unwrap(),
962                DynProofExpr::try_new_inequality(col_expr, high_expr, false).unwrap(),
963            )
964            .unwrap(),
965        )
966        .unwrap();
967
968        assert_eq!(expr_to_proof_expr(&expr, &schema).unwrap(), expected);
969    }
970
971    #[test]
972    fn we_can_convert_not_between_expr_to_proof_expr() {
973        let schema = vec![("column1".into(), ColumnType::BigInt)];
974
975        let col = df_column("namespace.table_name", "column1");
976        let low = Expr::Literal(ScalarValue::Int64(Some(10)));
977        let high = Expr::Literal(ScalarValue::Int64(Some(20)));
978        let expr = col.not_between(low, high);
979
980        let col_expr = DynProofExpr::new_column(ColumnRef::new(
981            TableRef::from_names(Some("namespace"), "table_name"),
982            "column1".into(),
983            ColumnType::BigInt,
984        ));
985        let low_expr = DynProofExpr::new_literal(LiteralValue::BigInt(10));
986        let high_expr = DynProofExpr::new_literal(LiteralValue::BigInt(20));
987
988        let expected = DynProofExpr::try_new_or(
989            DynProofExpr::try_new_inequality(col_expr.clone(), low_expr, true).unwrap(),
990            DynProofExpr::try_new_inequality(col_expr, high_expr, false).unwrap(),
991        )
992        .unwrap();
993
994        assert_eq!(expr_to_proof_expr(&expr, &schema).unwrap(), expected);
995    }
996
997    #[test]
998    fn we_can_extract_column_idents_from_between_expr() {
999        let col = df_column("table", "val");
1000        let low = Expr::Literal(ScalarValue::Int64(Some(1)));
1001        let high = Expr::Literal(ScalarValue::Int64(Some(100)));
1002        let expr = col.between(low, high);
1003        let result = get_column_idents_from_expr(&expr);
1004        let expected: IndexSet<Ident> = ["val".into()].into_iter().collect();
1005        assert_eq!(result, expected);
1006    }
1007
1008    #[test]
1009    fn we_can_get_proof_expr_for_timestamps_of_different_scale() {
1010        let lhs = Expr::Literal(ScalarValue::TimestampSecond(Some(1), None));
1011        let rhs = Expr::Literal(ScalarValue::TimestampNanosecond(Some(1), None));
1012        binary_expr_to_proof_expr(&lhs, &rhs, Operator::Gt, &Vec::new()).unwrap();
1013    }
1014
1015    // get_column_idents_from_expr tests
1016    #[test]
1017    fn we_can_extract_single_column_ident() {
1018        let expr = df_column("table", "column_a");
1019        let result = get_column_idents_from_expr(&expr);
1020        let expected: IndexSet<Ident> = ["column_a".into()].into_iter().collect();
1021        assert_eq!(result, expected);
1022    }
1023
1024    #[test]
1025    fn we_can_extract_column_idents_from_binary_expr() {
1026        let expr = df_column("table", "a").add(df_column("table", "b"));
1027        let result = get_column_idents_from_expr(&expr);
1028        let expected: IndexSet<Ident> = ["a".into(), "b".into()].into_iter().collect();
1029        assert_eq!(result, expected);
1030    }
1031
1032    #[test]
1033    fn we_can_extract_column_idents_from_nested_binary_expr() {
1034        // (a + b) * c
1035        let expr = df_column("table", "a")
1036            .add(df_column("table", "b"))
1037            .mul(df_column("table", "c"));
1038        let result = get_column_idents_from_expr(&expr);
1039        let expected: IndexSet<Ident> = ["a".into(), "b".into(), "c".into()].into_iter().collect();
1040        assert_eq!(result, expected);
1041    }
1042
1043    #[test]
1044    fn we_can_extract_column_idents_from_not_expr() {
1045        let expr = Expr::Not(Box::new(df_column("table", "bool_col")));
1046        let result = get_column_idents_from_expr(&expr);
1047        let expected: IndexSet<Ident> = ["bool_col".into()].into_iter().collect();
1048        assert_eq!(result, expected);
1049    }
1050
1051    #[test]
1052    fn we_can_extract_column_idents_from_alias_expr() {
1053        let expr = df_column("table", "col_x").alias("alias_name");
1054        let result = get_column_idents_from_expr(&expr);
1055        let expected: IndexSet<Ident> = ["col_x".into()].into_iter().collect();
1056        assert_eq!(result, expected);
1057    }
1058
1059    #[test]
1060    fn we_can_extract_column_idents_from_cast_expr() {
1061        let expr = Expr::Cast(Cast::new(
1062            Box::new(df_column("table", "num_col")),
1063            DataType::Int64,
1064        ));
1065        let result = get_column_idents_from_expr(&expr);
1066        let expected: IndexSet<Ident> = ["num_col".into()].into_iter().collect();
1067        assert_eq!(result, expected);
1068    }
1069
1070    #[test]
1071    fn we_can_extract_column_idents_from_aggregate_function() {
1072        let expr = Expr::AggregateFunction(datafusion::logical_expr::expr::AggregateFunction {
1073            func_def: datafusion::logical_expr::expr::AggregateFunctionDefinition::BuiltIn(
1074                datafusion::physical_plan::aggregates::AggregateFunction::Sum,
1075            ),
1076            args: vec![df_column("table", "value")],
1077            distinct: false,
1078            filter: None,
1079            order_by: None,
1080            null_treatment: None,
1081        });
1082        let result = get_column_idents_from_expr(&expr);
1083        let expected: IndexSet<Ident> = ["value".into()].into_iter().collect();
1084        assert_eq!(result, expected);
1085    }
1086
1087    #[test]
1088    fn we_can_extract_column_idents_from_aggregate_function_with_multiple_args() {
1089        let expr = Expr::AggregateFunction(datafusion::logical_expr::expr::AggregateFunction {
1090            func_def: datafusion::logical_expr::expr::AggregateFunctionDefinition::BuiltIn(
1091                datafusion::physical_plan::aggregates::AggregateFunction::Sum,
1092            ),
1093            args: vec![
1094                df_column("table", "col1"),
1095                df_column("table", "col2"),
1096                df_column("table", "col3"),
1097            ],
1098            distinct: false,
1099            filter: None,
1100            order_by: None,
1101            null_treatment: None,
1102        });
1103        let result = get_column_idents_from_expr(&expr);
1104        let expected: IndexSet<Ident> = ["col1".into(), "col2".into(), "col3".into()]
1105            .into_iter()
1106            .collect();
1107        assert_eq!(result, expected);
1108    }
1109
1110    #[test]
1111    fn we_can_extract_no_column_idents_from_literal() {
1112        let expr = Expr::Literal(ScalarValue::Int32(Some(42)));
1113        let result = get_column_idents_from_expr(&expr);
1114        assert!(result.is_empty());
1115    }
1116
1117    #[test]
1118    fn we_can_extract_column_idents_from_complex_nested_expr() {
1119        // NOT (a > b AND c < d)
1120        let inner = df_column("table", "a")
1121            .gt(df_column("table", "b"))
1122            .and(df_column("table", "c").lt(df_column("table", "d")));
1123        let expr = Expr::Not(Box::new(inner));
1124        let result = get_column_idents_from_expr(&expr);
1125        let expected: IndexSet<Ident> = ["a".into(), "b".into(), "c".into(), "d".into()]
1126            .into_iter()
1127            .collect();
1128        assert_eq!(result, expected);
1129    }
1130
1131    #[test]
1132    fn we_can_extract_column_idents_from_in_list_expr() {
1133        // a IN (b, c) references the needle column and every column in the list.
1134        let expr = df_column("table", "a").in_list(
1135            vec![df_column("table", "b"), df_column("table", "c")],
1136            false,
1137        );
1138        let result = get_column_idents_from_expr(&expr);
1139        let expected: IndexSet<Ident> = ["a".into(), "b".into(), "c".into()].into_iter().collect();
1140        assert_eq!(result, expected);
1141    }
1142
1143    #[test]
1144    fn we_can_extract_column_idents_preserving_order() {
1145        // IndexSet should preserve insertion order
1146        let expr = df_column("table", "z")
1147            .add(df_column("table", "a"))
1148            .add(df_column("table", "m"));
1149        let result = get_column_idents_from_expr(&expr);
1150        let idents: Vec<Ident> = result.into_iter().collect();
1151        assert_eq!(idents, vec!["z".into(), "a".into(), "m".into()]);
1152    }
1153
1154    #[test]
1155    fn we_can_handle_duplicate_column_references() {
1156        // a + a should only have 'a' once
1157        let expr = df_column("table", "a").add(df_column("table", "a"));
1158        let result = get_column_idents_from_expr(&expr);
1159        let expected: IndexSet<Ident> = ["a".into()].into_iter().collect();
1160        assert_eq!(result, expected);
1161    }
1162
1163    #[test]
1164    fn we_can_extract_columns_from_comparison_operations() {
1165        let expr = df_column("table", "price")
1166            .gt(df_column("table", "threshold"))
1167            .and(df_column("table", "active").eq(Expr::Literal(ScalarValue::Boolean(Some(true)))));
1168        let result = get_column_idents_from_expr(&expr);
1169        let expected: IndexSet<Ident> = ["price".into(), "threshold".into(), "active".into()]
1170            .into_iter()
1171            .collect();
1172        assert_eq!(result, expected);
1173    }
1174}