Skip to main content

gluesql_core/ast/
expr.rs

1use {
2    super::{
3        Aggregate, BinaryOperator, DataType, DateTimeField, Function, Literal, Query, ToSql,
4        ToSqlUnquoted, UnaryOperator,
5    },
6    crate::data::Value,
7    serde::{Deserialize, Serialize},
8    std::fmt::Write,
9};
10
11#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
12pub enum Expr {
13    Identifier(String),
14    CompoundIdentifier {
15        alias: String,
16        ident: String,
17    },
18    IsNull(Box<Expr>),
19    IsNotNull(Box<Expr>),
20    InList {
21        expr: Box<Expr>,
22        list: Vec<Expr>,
23        negated: bool,
24    },
25    InSubquery {
26        expr: Box<Expr>,
27        subquery: Box<Query>,
28        negated: bool,
29    },
30    Between {
31        expr: Box<Expr>,
32        negated: bool,
33        low: Box<Expr>,
34        high: Box<Expr>,
35    },
36    Like {
37        expr: Box<Expr>,
38        negated: bool,
39        pattern: Box<Expr>,
40    },
41    ILike {
42        expr: Box<Expr>,
43        negated: bool,
44        pattern: Box<Expr>,
45    },
46    Regex {
47        expr: Box<Expr>,
48        negated: bool,
49        pattern: Box<Expr>,
50        case_sensitive: bool,
51    },
52    BinaryOp {
53        left: Box<Expr>,
54        op: BinaryOperator,
55        right: Box<Expr>,
56    },
57    UnaryOp {
58        op: UnaryOperator,
59        expr: Box<Expr>,
60    },
61    Nested(Box<Expr>),
62    Literal(Literal),
63    Value(Value),
64    TypedString {
65        data_type: DataType,
66        value: String,
67    },
68    Function(Box<Function>),
69    Aggregate(Box<Aggregate>),
70    Exists {
71        subquery: Box<Query>,
72        negated: bool,
73    },
74    Subquery(Box<Query>),
75    Case {
76        operand: Option<Box<Expr>>,
77        when_then: Vec<(Expr, Expr)>,
78        else_result: Option<Box<Expr>>,
79    },
80    ArrayIndex {
81        obj: Box<Expr>,
82        indexes: Vec<Expr>,
83    },
84    Interval {
85        expr: Box<Expr>,
86        leading_field: Option<DateTimeField>,
87        last_field: Option<DateTimeField>,
88    },
89    Array {
90        elem: Vec<Expr>,
91    },
92}
93
94impl ToSql for Expr {
95    fn to_sql(&self) -> String {
96        self.to_sql_with(true)
97    }
98}
99
100impl ToSqlUnquoted for Expr {
101    fn to_sql_unquoted(&self) -> String {
102        self.to_sql_with(false)
103    }
104}
105
106impl Expr {
107    fn to_sql_with(&self, quoted: bool) -> String {
108        match self {
109            Expr::Identifier(s) => {
110                if quoted {
111                    format! {r#""{s}""#}
112                } else {
113                    s.to_owned()
114                }
115            }
116            Expr::BinaryOp { left, op, right } => {
117                format!(
118                    "{} {} {}",
119                    left.to_sql_with(quoted),
120                    op.to_sql(),
121                    right.to_sql_with(quoted),
122                )
123            }
124            Expr::CompoundIdentifier { alias, ident } => {
125                if quoted {
126                    format!(r#""{alias}"."{ident}""#)
127                } else {
128                    format!("{alias}.{ident}")
129                }
130            }
131            Expr::IsNull(s) => format!("{} IS NULL", s.to_sql_with(quoted)),
132            Expr::IsNotNull(s) => format!("{} IS NOT NULL", s.to_sql_with(quoted)),
133            Expr::InList {
134                expr,
135                list,
136                negated,
137            } => {
138                let expr = expr.to_sql_with(quoted);
139                let list = list
140                    .iter()
141                    .map(|expr| expr.to_sql_with(quoted))
142                    .collect::<Vec<_>>()
143                    .join(", ");
144
145                match negated {
146                    true => format!("{expr} NOT IN ({list})"),
147                    false => format!("{expr} IN ({list})"),
148                }
149            }
150            Expr::Between {
151                expr,
152                negated,
153                low,
154                high,
155            } => {
156                let expr = expr.to_sql_with(quoted);
157                let low = low.to_sql_with(quoted);
158                let high = high.to_sql_with(quoted);
159
160                match negated {
161                    true => format!("{expr} NOT BETWEEN {low} AND {high}"),
162                    false => format!("{expr} BETWEEN {low} AND {high}"),
163                }
164            }
165            Expr::Like {
166                expr,
167                negated,
168                pattern,
169            } => {
170                let expr = expr.to_sql_with(quoted);
171                let pattern = pattern.to_sql_with(quoted);
172
173                match negated {
174                    true => format!("{expr} NOT LIKE {pattern}"),
175                    false => format!("{expr} LIKE {pattern}"),
176                }
177            }
178            Expr::ILike {
179                expr,
180                negated,
181                pattern,
182            } => {
183                let expr = expr.to_sql_with(quoted);
184                let pattern = pattern.to_sql_with(quoted);
185
186                match negated {
187                    true => format!("{expr} NOT ILIKE {pattern}"),
188                    false => format!("{expr} ILIKE {pattern}"),
189                }
190            }
191            Expr::Regex {
192                expr,
193                negated,
194                pattern,
195                case_sensitive,
196            } => {
197                let op = match (*negated, *case_sensitive) {
198                    (false, true) => "~",
199                    (false, false) => "~*",
200                    (true, true) => "!~",
201                    (true, false) => "!~*",
202                };
203
204                format!(
205                    "{} {op} {}",
206                    expr.to_sql_with(quoted),
207                    pattern.to_sql_with(quoted)
208                )
209            }
210            Expr::UnaryOp { op, expr } => match op {
211                UnaryOperator::Factorial => {
212                    format!("{}{}", expr.to_sql_with(quoted), op.to_sql())
213                }
214                _ => format!("{}{}", op.to_sql(), expr.to_sql_with(quoted)),
215            },
216            Expr::Nested(expr) => format!("({})", expr.to_sql_with(quoted)),
217            Expr::Literal(s) => s.to_sql(),
218            Expr::Value(v) => v.to_sql(),
219            Expr::TypedString { data_type, value } => format!("{data_type} '{value}'"),
220            Expr::Case {
221                operand,
222                when_then,
223                else_result,
224            } => {
225                let operand = match operand {
226                    Some(operand) => format!("CASE {}", operand.to_sql_with(quoted)),
227                    None => "CASE".to_owned(),
228                };
229
230                let when_then = when_then
231                    .iter()
232                    .map(|(when, then)| {
233                        format!(
234                            "WHEN {} THEN {}",
235                            when.to_sql_with(quoted),
236                            then.to_sql_with(quoted)
237                        )
238                    })
239                    .collect::<Vec<_>>()
240                    .join("\n");
241
242                let else_result = else_result
243                    .as_ref()
244                    .map(|else_result| format!("ELSE {}", else_result.to_sql_with(quoted)));
245
246                match else_result {
247                    Some(else_result) => {
248                        [operand, when_then, else_result, "END".to_owned()].join("\n")
249                    }
250                    None => [operand, when_then, "END".to_owned()].join("\n"),
251                }
252            }
253            Expr::Aggregate(a) => a.to_sql(),
254            Expr::Function(func) => func.to_sql(),
255            Expr::InSubquery {
256                expr,
257                subquery,
258                negated,
259            } => match negated {
260                true => format!(
261                    "{} NOT IN ({})",
262                    expr.to_sql_with(quoted),
263                    subquery.to_sql()
264                ),
265                false => format!("{} IN ({})", expr.to_sql_with(quoted), subquery.to_sql()),
266            },
267            Expr::Exists { subquery, negated } => match negated {
268                true => format!("NOT EXISTS({})", subquery.to_sql()),
269                false => format!("EXISTS({})", subquery.to_sql()),
270            },
271            Expr::ArrayIndex { obj, indexes } => {
272                let obj = obj.to_sql_with(quoted);
273                let indexes = indexes.iter().fold(String::new(), |mut acc, index| {
274                    let _ = write!(acc, "[{}]", index.to_sql_with(quoted));
275                    acc
276                });
277                format!("{obj}{indexes}")
278            }
279            Expr::Array { elem } => {
280                let elem = elem
281                    .iter()
282                    .map(|e| e.to_sql_with(quoted))
283                    .collect::<Vec<_>>()
284                    .join(", ");
285                format!("[{elem}]")
286            }
287            Expr::Subquery(query) => format!("({})", query.to_sql()),
288            Expr::Interval {
289                expr,
290                leading_field,
291                last_field,
292            } => {
293                let expr = expr.to_sql_with(quoted);
294                let leading_field = leading_field
295                    .as_ref()
296                    .map_or_else(String::new, ToString::to_string);
297
298                match last_field {
299                    Some(last_field) => format!("INTERVAL {expr} {leading_field} TO {last_field}"),
300                    None => format!("INTERVAL {expr} {leading_field}"),
301                }
302            }
303        }
304    }
305}
306
307#[cfg(test)]
308mod tests {
309
310    use {
311        crate::ast::{
312            BinaryOperator, DataType, DateTimeField, Expr, Literal, Projection, Query, Select,
313            SelectItem, SetExpr, TableFactor, TableWithJoins, ToSql, ToSqlUnquoted, UnaryOperator,
314        },
315        bigdecimal::BigDecimal,
316        regex::Regex,
317        std::str::FromStr,
318    };
319
320    #[test]
321    fn to_sql() {
322        let re = Regex::new(r"\n\s+").unwrap();
323        let trim = |s: &str| re.replace_all(s.trim(), "\n").into_owned();
324
325        assert_eq!(r#""id""#, Expr::Identifier("id".to_owned()).to_sql());
326
327        assert_eq!(
328            r#""id" + "num""#,
329            Expr::BinaryOp {
330                left: Box::new(Expr::Identifier("id".to_owned())),
331                op: BinaryOperator::Plus,
332                right: Box::new(Expr::Identifier("num".to_owned()))
333            }
334            .to_sql()
335        );
336
337        for (negated, case_sensitive, expected) in [
338            (false, true, r#""id" ~ 'abc'"#),
339            (false, false, r#""id" ~* 'abc'"#),
340            (true, true, r#""id" !~ 'abc'"#),
341            (true, false, r#""id" !~* 'abc'"#),
342        ] {
343            assert_eq!(
344                Expr::Regex {
345                    expr: Box::new(Expr::Identifier("id".to_owned())),
346                    negated,
347                    pattern: Box::new(Expr::Literal(Literal::QuotedString("abc".to_owned()))),
348                    case_sensitive,
349                }
350                .to_sql(),
351                expected,
352            );
353        }
354        assert_eq!(
355            r#"-"id""#,
356            Expr::UnaryOp {
357                op: UnaryOperator::Minus,
358                expr: Box::new(Expr::Identifier("id".to_owned())),
359            }
360            .to_sql(),
361        );
362
363        assert_eq!(
364            r#""alias"."column""#,
365            Expr::CompoundIdentifier {
366                alias: "alias".into(),
367                ident: "column".into()
368            }
369            .to_sql()
370        );
371
372        assert_eq!(
373            "alias.column",
374            Expr::CompoundIdentifier {
375                alias: "alias".into(),
376                ident: "column".into()
377            }
378            .to_sql_unquoted()
379        );
380
381        let id_expr: Box<Expr> = Box::new(Expr::Identifier("id".to_owned()));
382        assert_eq!(r#""id" IS NULL"#, Expr::IsNull(id_expr).to_sql());
383
384        let id_expr: Box<Expr> = Box::new(Expr::Identifier("id".to_owned()));
385        assert_eq!(r#""id" IS NOT NULL"#, Expr::IsNotNull(id_expr).to_sql());
386
387        assert_eq!(
388            "INT '1'",
389            Expr::TypedString {
390                data_type: DataType::Int,
391                value: "1".to_owned()
392            }
393            .to_sql()
394        );
395
396        assert_eq!(
397            r#"("id")"#,
398            Expr::Nested(Box::new(Expr::Identifier("id".to_owned()))).to_sql(),
399        );
400
401        assert_eq!(
402            r#""id" BETWEEN "low" AND "high""#,
403            Expr::Between {
404                expr: Box::new(Expr::Identifier("id".to_owned())),
405                negated: false,
406                low: Box::new(Expr::Identifier("low".to_owned())),
407                high: Box::new(Expr::Identifier("high".to_owned()))
408            }
409            .to_sql()
410        );
411
412        assert_eq!(
413            r#""id" NOT BETWEEN "low" AND "high""#,
414            Expr::Between {
415                expr: Box::new(Expr::Identifier("id".to_owned())),
416                negated: true,
417                low: Box::new(Expr::Identifier("low".to_owned())),
418                high: Box::new(Expr::Identifier("high".to_owned()))
419            }
420            .to_sql()
421        );
422
423        assert_eq!(
424            r#""id" LIKE '%abc'"#,
425            Expr::Like {
426                expr: Box::new(Expr::Identifier("id".to_owned())),
427                negated: false,
428                pattern: Box::new(Expr::Literal(Literal::QuotedString("%abc".to_owned()))),
429            }
430            .to_sql()
431        );
432
433        assert_eq!(
434            r#""id" NOT LIKE '%abc'"#,
435            Expr::Like {
436                expr: Box::new(Expr::Identifier("id".to_owned())),
437                negated: true,
438                pattern: Box::new(Expr::Literal(Literal::QuotedString("%abc".to_owned()))),
439            }
440            .to_sql()
441        );
442
443        assert_eq!(
444            r#""id" ILIKE '%abc_'"#,
445            Expr::ILike {
446                expr: Box::new(Expr::Identifier("id".to_owned())),
447                negated: false,
448                pattern: Box::new(Expr::Literal(Literal::QuotedString("%abc_".to_owned()))),
449            }
450            .to_sql()
451        );
452
453        assert_eq!(
454            r#""id" NOT ILIKE '%abc_'"#,
455            Expr::ILike {
456                expr: Box::new(Expr::Identifier("id".to_owned())),
457                negated: true,
458                pattern: Box::new(Expr::Literal(Literal::QuotedString("%abc_".to_owned()))),
459            }
460            .to_sql()
461        );
462
463        assert_eq!(
464            r#""id" IN ('a', 'b', 'c')"#,
465            Expr::InList {
466                expr: Box::new(Expr::Identifier("id".to_owned())),
467                list: vec![
468                    Expr::Literal(Literal::QuotedString("a".to_owned())),
469                    Expr::Literal(Literal::QuotedString("b".to_owned())),
470                    Expr::Literal(Literal::QuotedString("c".to_owned()))
471                ],
472                negated: false
473            }
474            .to_sql()
475        );
476
477        assert_eq!(
478            r#""id" NOT IN ('a', 'b', 'c')"#,
479            Expr::InList {
480                expr: Box::new(Expr::Identifier("id".to_owned())),
481                list: vec![
482                    Expr::Literal(Literal::QuotedString("a".to_owned())),
483                    Expr::Literal(Literal::QuotedString("b".to_owned())),
484                    Expr::Literal(Literal::QuotedString("c".to_owned()))
485                ],
486                negated: true
487            }
488            .to_sql()
489        );
490
491        assert_eq!(
492            r#""id" IN (SELECT * FROM "FOO")"#,
493            Expr::InSubquery {
494                expr: Box::new(Expr::Identifier("id".to_owned())),
495                subquery: Box::new(Query {
496                    body: SetExpr::Select(Box::new(Select {
497                        distinct: false,
498                        projection: Projection::SelectItems(vec![SelectItem::Wildcard]),
499                        from: TableWithJoins {
500                            relation: TableFactor::Table {
501                                name: "FOO".to_owned(),
502                                alias: None,
503                            },
504                            joins: Vec::new(),
505                        },
506                        selection: None,
507                        group_by: Vec::new(),
508                        having: None,
509                    })),
510                    order_by: Vec::new(),
511                    limit: None,
512                    offset: None,
513                }),
514                negated: false
515            }
516            .to_sql()
517        );
518
519        assert_eq!(
520            r#""id" NOT IN (SELECT * FROM "FOO")"#,
521            Expr::InSubquery {
522                expr: Box::new(Expr::Identifier("id".to_owned())),
523                subquery: Box::new(Query {
524                    body: SetExpr::Select(Box::new(Select {
525                        distinct: false,
526                        projection: Projection::SelectItems(vec![SelectItem::Wildcard]),
527                        from: TableWithJoins {
528                            relation: TableFactor::Table {
529                                name: "FOO".to_owned(),
530                                alias: None,
531                            },
532                            joins: Vec::new(),
533                        },
534                        selection: None,
535                        group_by: Vec::new(),
536                        having: None,
537                    })),
538                    order_by: Vec::new(),
539                    limit: None,
540                    offset: None,
541                }),
542                negated: true
543            }
544            .to_sql()
545        );
546
547        assert_eq!(
548            r#"EXISTS(SELECT * FROM "FOO")"#,
549            Expr::Exists {
550                subquery: Box::new(Query {
551                    body: SetExpr::Select(Box::new(Select {
552                        distinct: false,
553                        projection: Projection::SelectItems(vec![SelectItem::Wildcard]),
554                        from: TableWithJoins {
555                            relation: TableFactor::Table {
556                                name: "FOO".to_owned(),
557                                alias: None,
558                            },
559                            joins: Vec::new(),
560                        },
561                        selection: None,
562                        group_by: Vec::new(),
563                        having: None,
564                    })),
565                    order_by: Vec::new(),
566                    limit: None,
567                    offset: None,
568                }),
569                negated: false,
570            }
571            .to_sql(),
572        );
573
574        assert_eq!(
575            r#"NOT EXISTS(SELECT * FROM "FOO")"#,
576            Expr::Exists {
577                subquery: Box::new(Query {
578                    body: SetExpr::Select(Box::new(Select {
579                        distinct: false,
580                        projection: Projection::SelectItems(vec![SelectItem::Wildcard]),
581                        from: TableWithJoins {
582                            relation: TableFactor::Table {
583                                name: "FOO".to_owned(),
584                                alias: None,
585                            },
586                            joins: Vec::new(),
587                        },
588                        selection: None,
589                        group_by: Vec::new(),
590                        having: None,
591                    })),
592                    order_by: Vec::new(),
593                    limit: None,
594                    offset: None,
595                }),
596                negated: true,
597            }
598            .to_sql(),
599        );
600
601        assert_eq!(
602            r#"(SELECT * FROM "FOO")"#,
603            Expr::Subquery(Box::new(Query {
604                body: SetExpr::Select(Box::new(Select {
605                    distinct: false,
606                    projection: Projection::SelectItems(vec![SelectItem::Wildcard]),
607                    from: TableWithJoins {
608                        relation: TableFactor::Table {
609                            name: "FOO".to_owned(),
610                            alias: None,
611                        },
612                        joins: Vec::new(),
613                    },
614                    selection: None,
615                    group_by: Vec::new(),
616                    having: None,
617                })),
618                order_by: Vec::new(),
619                limit: None,
620                offset: None,
621            }))
622            .to_sql()
623        );
624
625        assert_eq!(
626            trim(
627                r#"CASE "id"
628                  WHEN 1 THEN 'a'
629                  WHEN 2 THEN 'b'
630                  ELSE 'c'
631                END"#,
632            ),
633            Expr::Case {
634                operand: Some(Box::new(Expr::Identifier("id".to_owned()))),
635                when_then: vec![
636                    (
637                        Expr::Literal(Literal::Number(BigDecimal::from_str("1").unwrap())),
638                        Expr::Literal(Literal::QuotedString("a".to_owned()))
639                    ),
640                    (
641                        Expr::Literal(Literal::Number(BigDecimal::from_str("2").unwrap())),
642                        Expr::Literal(Literal::QuotedString("b".to_owned()))
643                    )
644                ],
645                else_result: Some(Box::new(Expr::Literal(Literal::QuotedString(
646                    "c".to_owned()
647                ))))
648            }
649            .to_sql()
650        );
651
652        assert_eq!(
653            trim(
654                r#"CASE
655                  WHEN "id" = 1 THEN 'a'
656                  WHEN "id" = 2 THEN 'b'
657                END"#,
658            ),
659            Expr::Case {
660                operand: None,
661                when_then: vec![
662                    (
663                        Expr::BinaryOp {
664                            left: Box::new(Expr::Identifier("id".to_owned())),
665                            op: BinaryOperator::Eq,
666                            right: Box::new(Expr::Literal(Literal::Number(
667                                BigDecimal::from_str("1").unwrap()
668                            )))
669                        },
670                        Expr::Literal(Literal::QuotedString("a".to_owned()))
671                    ),
672                    (
673                        Expr::BinaryOp {
674                            left: Box::new(Expr::Identifier("id".to_owned())),
675                            op: BinaryOperator::Eq,
676                            right: Box::new(Expr::Literal(Literal::Number(
677                                BigDecimal::from_str("2").unwrap()
678                            )))
679                        },
680                        Expr::Literal(Literal::QuotedString("b".to_owned()))
681                    )
682                ],
683                else_result: None,
684            }
685            .to_sql()
686        );
687
688        assert_eq!(
689            trim(
690                r#"CASE "id"
691                  WHEN 1 THEN 'a'
692                  WHEN 2 THEN 'b'
693                END"#,
694            ),
695            Expr::Case {
696                operand: Some(Box::new(Expr::Identifier("id".to_owned()))),
697                when_then: vec![
698                    (
699                        Expr::Literal(Literal::Number(BigDecimal::from_str("1").unwrap())),
700                        Expr::Literal(Literal::QuotedString("a".to_owned()))
701                    ),
702                    (
703                        Expr::Literal(Literal::Number(BigDecimal::from_str("2").unwrap())),
704                        Expr::Literal(Literal::QuotedString("b".to_owned()))
705                    )
706                ],
707                else_result: None,
708            }
709            .to_sql()
710        );
711
712        assert_eq!(
713            r#""choco"[1][2]"#,
714            Expr::ArrayIndex {
715                obj: Box::new(Expr::Identifier("choco".to_owned())),
716                indexes: vec![
717                    Expr::Literal(Literal::Number(BigDecimal::from_str("1").unwrap())),
718                    Expr::Literal(Literal::Number(BigDecimal::from_str("2").unwrap()))
719                ]
720            }
721            .to_sql()
722        );
723
724        assert_eq!(
725            r"['GlueSQL', 'Rust']",
726            Expr::Array {
727                elem: vec![
728                    Expr::Literal(Literal::QuotedString("GlueSQL".to_owned())),
729                    Expr::Literal(Literal::QuotedString("Rust".to_owned()))
730                ]
731            }
732            .to_sql()
733        );
734
735        assert_eq!(
736            r#"INTERVAL "col1" + 3 DAY"#,
737            &Expr::Interval {
738                expr: Box::new(Expr::BinaryOp {
739                    left: Box::new(Expr::Identifier("col1".to_owned())),
740                    op: BinaryOperator::Plus,
741                    right: Box::new(Expr::Literal(Literal::Number(3.into()))),
742                }),
743                leading_field: Some(DateTimeField::Day),
744                last_field: None,
745            }
746            .to_sql()
747        );
748
749        assert_eq!(
750            "INTERVAL '3-5' HOUR TO MINUTE",
751            &Expr::Interval {
752                expr: Box::new(Expr::Literal(Literal::QuotedString("3-5".to_owned()))),
753                leading_field: Some(DateTimeField::Hour),
754                last_field: Some(DateTimeField::Minute),
755            }
756            .to_sql()
757        );
758    }
759}