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
19pub(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
60fn 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#[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
150pub 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 let needle = expr_to_proof_expr(expr, schema)?;
187 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 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 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 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
271fn 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 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 #[test]
382 fn we_can_convert_alias_to_proof_expr() {
383 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 #[test]
391 fn we_can_convert_column_expr_to_proof_expr() {
392 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 #[test]
400 fn we_can_convert_in_list_to_proof_expr() {
401 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 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 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 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 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 #[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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 #[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 #[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 #[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 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 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 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 #[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 #[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 #[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 #[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 #[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 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 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 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 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 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}