Skip to main content

drizzle_core/expr/
ops.rs

1//! Arithmetic on [`SQLExpr`] with the Rust operators `+ - * / %` and unary `-`.
2//!
3//! Both operands must be numeric. The result type follows the operand types
4//! (see [`ArithmeticOutput`](crate::types::ArithmeticOutput)). The result is
5//! nullable if either operand is, and also for `/` and `%` on SQLite and
6//! MySQL, where dividing by zero gives NULL. It is an aggregate if either
7//! operand is. Operands are parenthesized when needed to keep the grouping of
8//! the Rust expression, so `a * (b + c)` renders as written.
9
10use core::ops::{Add, Div, Mul, Neg, Rem, Sub};
11
12use crate::dialect::Dialect;
13use crate::sql::{SQL, SQLChunk, Token};
14use crate::traits::SQLParam;
15use crate::types::{AddOp, ArithmeticOutput, DivOp, MulOp, NegOutput, Numeric, RemOp, SubOp};
16
17use super::{AggregateKind, Expr, Nullability, ResolveArithmeticNullability, SQLExpr};
18
19type ArithmeticNullable<'a, V, T, N, Rhs, Op> = <<T as ArithmeticOutput<
20    <Rhs as Expr<'a, V>>::SQLType,
21    Op,
22>>::Nullability as ResolveArithmeticNullability<
23    N,
24    <Rhs as Expr<'a, V>>::Nullable,
25>>::Output;
26
27#[inline]
28fn binary_op_sql<'a, V, L, R>(left: L, operator: Token, right: R) -> SQL<'a, V>
29where
30    V: SQLParam + 'a,
31    L: Expr<'a, V>,
32    R: Expr<'a, V>,
33{
34    binary_operator_sql(left.into_expr_sql(), operator, right.into_expr_sql())
35}
36
37// =============================================================================
38// Operator precedence
39// =============================================================================
40
41/// Rank given to anything between two operands that is not a ranked binary
42/// operator: comparisons, logical operators, raw operator text. It is below
43/// every ranked operator, so such an operand is always grouped.
44const LOOSEST: u8 = 0;
45
46/// How tightly `operator` binds between two operands in `dialect`; larger
47/// binds tighter. Only the operators the expression builders place between
48/// two operands are ranked.
49const fn binding_power(dialect: Dialect, operator: Token) -> u8 {
50    match dialect {
51        // SQLite ranks `||` above `*`, and `& | << >>` together below `+ -`.
52        Dialect::SQLite => match operator {
53            Token::CONCAT => 5,
54            Token::STAR | Token::SLASH | Token::REM => 4,
55            Token::PLUS | Token::MINUS => 3,
56            Token::BITAND | Token::BITOR | Token::LSHIFT | Token::RSHIFT => 2,
57            _ => LOOSEST,
58        },
59        // PostgreSQL puts `||` and the bitwise operators in its shared
60        // "any other operator" level, below `+ -`.
61        Dialect::PostgreSQL => match operator {
62            Token::STAR | Token::SLASH | Token::REM => 4,
63            Token::PLUS | Token::MINUS => 3,
64            Token::CONCAT | Token::BITAND | Token::BITOR | Token::LSHIFT | Token::RSHIFT => 2,
65            _ => LOOSEST,
66        },
67        // MySQL reads `||` as logical OR, so it stays unranked.
68        Dialect::MySQL => match operator {
69            Token::STAR | Token::SLASH | Token::REM => 6,
70            Token::PLUS | Token::MINUS => 5,
71            Token::LSHIFT | Token::RSHIFT => 4,
72            Token::BITAND => 3,
73            Token::BITOR => 2,
74            _ => LOOSEST,
75        },
76    }
77}
78
79/// The loosest-binding operator found at the top level of an operand.
80#[derive(Clone, Copy)]
81struct TopLevelOperator {
82    power: u8,
83    /// Every operator at `power` is `||`, which is associative.
84    only_concat: bool,
85}
86
87impl TopLevelOperator {
88    fn record(found: &mut Option<Self>, power: u8, concat: bool) {
89        match found {
90            None => {
91                *found = Some(Self {
92                    power,
93                    only_concat: concat,
94                });
95            }
96            Some(current) if power < current.power => {
97                *current = Self {
98                    power,
99                    only_concat: concat,
100                };
101            }
102            Some(current) if power == current.power => current.only_concat &= concat,
103            Some(_) => {}
104        }
105    }
106}
107
108/// Keyword and comparison tokens that join two operands at a looser level
109/// than any arithmetic operator.
110const fn is_loose_infix(token: Token) -> bool {
111    matches!(
112        token,
113        Token::EQ
114            | Token::NE
115            | Token::LT
116            | Token::GT
117            | Token::LE
118            | Token::GE
119            | Token::AND
120            | Token::OR
121            | Token::NOT
122            | Token::IS
123            | Token::ISNOT
124            | Token::IN
125            | Token::LIKE
126            | Token::BETWEEN
127            | Token::ESCAPE
128            | Token::ISNULL
129            | Token::NOTNULL
130            | Token::MATCH
131    )
132}
133
134/// Raw text that reads as a single term (a function name, keyword literal,
135/// or number) rather than as an operator.
136fn is_term_text(text: &str) -> bool {
137    let text = text.trim();
138    !text.is_empty()
139        && text
140            .chars()
141            .all(|ch| ch.is_alphanumeric() || matches!(ch, '_' | '.' | '"' | '`' | '\''))
142}
143
144/// Finds the loosest binary operator outside any parentheses or
145/// `CASE ... END` in `operand`. `None` means the operand is a single term: a
146/// column, value, function call, parenthesized group, or a sign-prefixed term.
147fn top_level_operator<V: SQLParam>(operand: &SQL<'_, V>) -> Option<TopLevelOperator> {
148    let mut depth = 0usize;
149    // True at the start and after an operator, where `-`/`+` are signs.
150    let mut expect_term = true;
151    let mut found = None;
152
153    for chunk in &operand.chunks {
154        match chunk {
155            SQLChunk::Token(Token::LPAREN | Token::CASE) => {
156                depth += 1;
157                expect_term = false;
158            }
159            SQLChunk::Token(Token::RPAREN | Token::END) => {
160                depth = depth.saturating_sub(1);
161                expect_term = false;
162            }
163            _ if depth > 0 => {}
164            SQLChunk::Token(
165                token @ (Token::PLUS
166                | Token::MINUS
167                | Token::STAR
168                | Token::SLASH
169                | Token::REM
170                | Token::CONCAT
171                | Token::BITAND
172                | Token::BITOR
173                | Token::LSHIFT
174                | Token::RSHIFT),
175            ) => {
176                if !expect_term {
177                    TopLevelOperator::record(
178                        &mut found,
179                        binding_power(V::DIALECT, *token),
180                        matches!(token, Token::CONCAT),
181                    );
182                    expect_term = true;
183                }
184            }
185            SQLChunk::Token(Token::BITNOT) => {}
186            SQLChunk::Token(token) if is_loose_infix(*token) => {
187                TopLevelOperator::record(&mut found, LOOSEST, false);
188                expect_term = true;
189            }
190            SQLChunk::Raw(text) if expect_term && matches!(text.trim(), "-" | "+") => {}
191            SQLChunk::Raw(text) if !is_term_text(text) => {
192                TopLevelOperator::record(&mut found, LOOSEST, false);
193                expect_term = true;
194            }
195            _ => expect_term = false,
196        }
197    }
198
199    found
200}
201
202/// Whether `operand` must be parenthesized to stay a single operand of
203/// `operator`. SQL operators of equal precedence associate to the left, so a
204/// right-hand operand also needs grouping at equal precedence (`a - (b - c)`
205/// is not `a - b - c`); the associative `||` is the exception.
206fn needs_grouping<V: SQLParam>(operand: &SQL<'_, V>, operator: Token, right_hand: bool) -> bool {
207    let Some(inner) = top_level_operator(operand) else {
208        return false;
209    };
210    let outer = binding_power(V::DIALECT, operator);
211    if right_hand {
212        inner.power < outer
213            || (inner.power == outer && !(matches!(operator, Token::CONCAT) && inner.only_concat))
214    } else {
215        inner.power < outer
216    }
217}
218
219/// Renders `left operator right` so the database evaluates the tree the Rust
220/// expression built.
221///
222/// An operand that is itself a binary expression is parenthesized when its
223/// top-level operator binds more loosely than `operator` (on the right-hand
224/// side, also when it binds equally), so `a * (b + c)` keeps its grouping.
225/// Operands that already read correctly stay flat: a single term renders as
226/// before, and so does a chain such as `a * b + c`.
227pub(crate) fn binary_operator_sql<'a, V>(
228    left: SQL<'a, V>,
229    operator: Token,
230    right: SQL<'a, V>,
231) -> SQL<'a, V>
232where
233    V: SQLParam + 'a,
234{
235    let left = left.parens_if_subquery();
236    let right = right.parens_if_subquery();
237    let left = if needs_grouping(&left, operator, false) {
238        left.parens()
239    } else {
240        left
241    };
242    let right = if needs_grouping(&right, operator, true) {
243        right.parens()
244    } else {
245        right
246    };
247    left.push(operator).append(right)
248}
249
250// =============================================================================
251// Addition
252// =============================================================================
253
254impl<'a, V, T, N, A, S, Rhs> Add<Rhs> for SQLExpr<'a, V, T, N, A, S>
255where
256    V: SQLParam + 'a,
257    T: ArithmeticOutput<Rhs::SQLType, AddOp>,
258    N: Nullability,
259    A: AggregateKind,
260    Rhs: Expr<'a, V>,
261    Rhs::SQLType: Numeric,
262    Rhs::Nullable: Nullability,
263    <T as ArithmeticOutput<Rhs::SQLType, AddOp>>::Nullability:
264        ResolveArithmeticNullability<N, Rhs::Nullable>,
265{
266    type Output = SQLExpr<
267        'a,
268        V,
269        <T as ArithmeticOutput<Rhs::SQLType, AddOp>>::Output,
270        ArithmeticNullable<'a, V, T, N, Rhs, AddOp>,
271        <A as AggregateKind>::Or<Rhs::Aggregate>,
272        (S, Rhs::Sources),
273    >;
274
275    fn add(self, rhs: Rhs) -> Self::Output {
276        SQLExpr::new(binary_op_sql(self, Token::PLUS, rhs))
277    }
278}
279
280// =============================================================================
281// Subtraction
282// =============================================================================
283
284impl<'a, V, T, N, A, S, Rhs> Sub<Rhs> for SQLExpr<'a, V, T, N, A, S>
285where
286    V: SQLParam + 'a,
287    T: ArithmeticOutput<Rhs::SQLType, SubOp>,
288    N: Nullability,
289    A: AggregateKind,
290    Rhs: Expr<'a, V>,
291    Rhs::SQLType: Numeric,
292    Rhs::Nullable: Nullability,
293    <T as ArithmeticOutput<Rhs::SQLType, SubOp>>::Nullability:
294        ResolveArithmeticNullability<N, Rhs::Nullable>,
295{
296    type Output = SQLExpr<
297        'a,
298        V,
299        <T as ArithmeticOutput<Rhs::SQLType, SubOp>>::Output,
300        ArithmeticNullable<'a, V, T, N, Rhs, SubOp>,
301        <A as AggregateKind>::Or<Rhs::Aggregate>,
302        (S, Rhs::Sources),
303    >;
304
305    fn sub(self, rhs: Rhs) -> Self::Output {
306        SQLExpr::new(binary_op_sql(self, Token::MINUS, rhs))
307    }
308}
309
310// =============================================================================
311// Multiplication
312// =============================================================================
313
314impl<'a, V, T, N, A, S, Rhs> Mul<Rhs> for SQLExpr<'a, V, T, N, A, S>
315where
316    V: SQLParam + 'a,
317    T: ArithmeticOutput<Rhs::SQLType, MulOp>,
318    N: Nullability,
319    A: AggregateKind,
320    Rhs: Expr<'a, V>,
321    Rhs::SQLType: Numeric,
322    Rhs::Nullable: Nullability,
323    <T as ArithmeticOutput<Rhs::SQLType, MulOp>>::Nullability:
324        ResolveArithmeticNullability<N, Rhs::Nullable>,
325{
326    type Output = SQLExpr<
327        'a,
328        V,
329        <T as ArithmeticOutput<Rhs::SQLType, MulOp>>::Output,
330        ArithmeticNullable<'a, V, T, N, Rhs, MulOp>,
331        <A as AggregateKind>::Or<Rhs::Aggregate>,
332        (S, Rhs::Sources),
333    >;
334
335    fn mul(self, rhs: Rhs) -> Self::Output {
336        SQLExpr::new(binary_op_sql(self, Token::STAR, rhs))
337    }
338}
339
340// =============================================================================
341// Division
342// =============================================================================
343
344impl<'a, V, T, N, A, S, Rhs> Div<Rhs> for SQLExpr<'a, V, T, N, A, S>
345where
346    V: SQLParam + 'a,
347    T: ArithmeticOutput<Rhs::SQLType, DivOp>,
348    N: Nullability,
349    A: AggregateKind,
350    Rhs: Expr<'a, V>,
351    Rhs::SQLType: Numeric,
352    Rhs::Nullable: Nullability,
353    <T as ArithmeticOutput<Rhs::SQLType, DivOp>>::Nullability:
354        ResolveArithmeticNullability<N, Rhs::Nullable>,
355{
356    type Output = SQLExpr<
357        'a,
358        V,
359        <T as ArithmeticOutput<Rhs::SQLType, DivOp>>::Output,
360        ArithmeticNullable<'a, V, T, N, Rhs, DivOp>,
361        <A as AggregateKind>::Or<Rhs::Aggregate>,
362        (S, Rhs::Sources),
363    >;
364
365    fn div(self, rhs: Rhs) -> Self::Output {
366        SQLExpr::new(binary_op_sql(self, Token::SLASH, rhs))
367    }
368}
369
370// =============================================================================
371// Remainder (Modulo)
372// =============================================================================
373
374impl<'a, V, T, N, A, S, Rhs> Rem<Rhs> for SQLExpr<'a, V, T, N, A, S>
375where
376    V: SQLParam + 'a,
377    T: ArithmeticOutput<Rhs::SQLType, RemOp>,
378    N: Nullability,
379    A: AggregateKind,
380    Rhs: Expr<'a, V>,
381    Rhs::SQLType: Numeric,
382    Rhs::Nullable: Nullability,
383    <T as ArithmeticOutput<Rhs::SQLType, RemOp>>::Nullability:
384        ResolveArithmeticNullability<N, Rhs::Nullable>,
385{
386    type Output = SQLExpr<
387        'a,
388        V,
389        <T as ArithmeticOutput<Rhs::SQLType, RemOp>>::Output,
390        ArithmeticNullable<'a, V, T, N, Rhs, RemOp>,
391        <A as AggregateKind>::Or<Rhs::Aggregate>,
392        (S, Rhs::Sources),
393    >;
394
395    fn rem(self, rhs: Rhs) -> Self::Output {
396        SQLExpr::new(binary_op_sql(self, Token::REM, rhs))
397    }
398}
399
400// =============================================================================
401// Negation
402// =============================================================================
403
404impl<'a, V, T, N, A, S> Neg for SQLExpr<'a, V, T, N, A, S>
405where
406    V: SQLParam + 'a,
407    T: Numeric + NegOutput,
408    N: Nullability,
409    A: AggregateKind,
410{
411    type Output = SQLExpr<'a, V, T::Output, N, A, S>;
412
413    fn neg(self) -> Self::Output {
414        SQLExpr::new(SQL::from(Token::MINUS).append(self.into_expr_sql().parens()))
415    }
416}
417
418#[cfg(test)]
419mod tests {
420    use super::binary_operator_sql;
421    use crate::sql::{SQL, Token};
422    use crate::{Dialect, MySQLDialect, PostgresDialect, SQLParam, SQLiteDialect};
423
424    #[derive(Clone, Debug)]
425    struct SqliteParam;
426
427    impl SQLParam for SqliteParam {
428        const DIALECT: Dialect = Dialect::SQLite;
429        type DialectMarker = SQLiteDialect;
430    }
431
432    #[derive(Clone, Debug)]
433    struct PostgresParam;
434
435    impl SQLParam for PostgresParam {
436        const DIALECT: Dialect = Dialect::PostgreSQL;
437        type DialectMarker = PostgresDialect;
438    }
439
440    #[derive(Clone, Debug)]
441    struct MySqlParam;
442
443    impl SQLParam for MySqlParam {
444        const DIALECT: Dialect = Dialect::MySQL;
445        type DialectMarker = MySQLDialect;
446    }
447
448    fn term<V: SQLParam>(name: &'static str) -> SQL<'static, V> {
449        SQL::ident(name)
450    }
451
452    fn apply<V: SQLParam>(
453        left: SQL<'static, V>,
454        operator: Token,
455        right: SQL<'static, V>,
456    ) -> SQL<'static, V> {
457        binary_operator_sql(left, operator, right)
458    }
459
460    #[test]
461    fn single_operator_stays_flat() {
462        let product = apply::<SqliteParam>(term("a"), Token::STAR, term("b"));
463        assert_eq!(product.sql(), r#""a" * "b""#);
464    }
465
466    #[test]
467    fn looser_operand_is_grouped_on_either_side() {
468        let sum = || apply::<SqliteParam>(term("b"), Token::PLUS, term("c"));
469        assert_eq!(
470            apply(term("a"), Token::STAR, sum()).sql(),
471            r#""a" *("b" + "c")"#
472        );
473        assert_eq!(
474            apply(sum(), Token::STAR, term("a")).sql(),
475            r#"("b" + "c")* "a""#
476        );
477    }
478
479    #[test]
480    fn tighter_left_chain_stays_flat() {
481        let product = apply::<PostgresParam>(term("a"), Token::STAR, term("b"));
482        let chain = apply(product, Token::PLUS, term("c"));
483        assert_eq!(chain.sql(), r#""a" * "b" + "c""#);
484
485        let difference = apply::<PostgresParam>(term("a"), Token::MINUS, term("b"));
486        let chain = apply(difference, Token::MINUS, term("c"));
487        assert_eq!(chain.sql(), r#""a" - "b" - "c""#);
488    }
489
490    #[test]
491    fn equal_precedence_on_the_right_is_grouped() {
492        let difference = apply::<MySqlParam>(term("b"), Token::MINUS, term("c"));
493        assert_eq!(
494            apply(term("a"), Token::MINUS, difference).sql(),
495            "`a` -(`b` - `c`)"
496        );
497    }
498
499    #[test]
500    fn concatenation_chain_stays_flat_on_the_right() {
501        let tail = apply::<SqliteParam>(term("b"), Token::CONCAT, term("c"));
502        assert_eq!(
503            apply(term("a"), Token::CONCAT, tail).sql(),
504            r#""a" || "b" || "c""#
505        );
506    }
507
508    #[test]
509    fn concatenation_precedence_follows_the_dialect() {
510        // SQLite binds `||` tighter than `+`; PostgreSQL binds it looser.
511        let sqlite_sum = apply::<SqliteParam>(term("b"), Token::PLUS, term("c"));
512        assert_eq!(
513            apply(term("a"), Token::CONCAT, sqlite_sum).sql(),
514            r#""a" ||("b" + "c")"#
515        );
516
517        let postgres_sum = apply::<PostgresParam>(term("b"), Token::PLUS, term("c"));
518        assert_eq!(
519            apply(term("a"), Token::CONCAT, postgres_sum).sql(),
520            r#""a" || "b" + "c""#
521        );
522    }
523
524    #[test]
525    fn terms_with_inner_operators_stay_flat() {
526        // A function call, a parenthesized group and a signed term are each
527        // a single operand, whatever they contain.
528        let call = SQL::<SqliteParam>::raw("ABS")
529            .push(Token::LPAREN)
530            .append(apply(term("b"), Token::MINUS, term("c")))
531            .push(Token::RPAREN);
532        assert_eq!(
533            apply(term("a"), Token::STAR, call).sql(),
534            r#""a" * ABS ("b" - "c")"#
535        );
536
537        let signed = SQL::<SqliteParam>::raw("-").append(term("b"));
538        assert_eq!(
539            apply(term("a"), Token::MINUS, signed).sql(),
540            r#""a" - - "b""#
541        );
542    }
543
544    #[test]
545    fn comparison_operand_is_grouped() {
546        let comparison = SQL::<SqliteParam>::ident("b")
547            .push(Token::EQ)
548            .append(term("c"));
549        assert_eq!(
550            apply(term("a"), Token::PLUS, comparison).sql(),
551            r#""a" +("b" = "c")"#
552        );
553    }
554}