1use 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
37const LOOSEST: u8 = 0;
45
46const fn binding_power(dialect: Dialect, operator: Token) -> u8 {
50 match dialect {
51 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 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 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#[derive(Clone, Copy)]
81struct TopLevelOperator {
82 power: u8,
83 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
108const 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
134fn 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
144fn top_level_operator<V: SQLParam>(operand: &SQL<'_, V>) -> Option<TopLevelOperator> {
148 let mut depth = 0usize;
149 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
202fn 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
219pub(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
250impl<'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
280impl<'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
310impl<'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
340impl<'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
370impl<'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
400impl<'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 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 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}