drizzle_core/expr/case.rs
1//! `CASE WHEN ... THEN ... ELSE ... END` expressions.
2//!
3//! [`case`] starts a builder. The first `.when(condition, result)` fixes the
4//! result type; later branches and the `ELSE` value must have a compatible
5//! type. Finish with [`r#else`](CaseBuilder#method.else) or
6//! [`end`](CaseBuilder::end).
7
8use core::marker::PhantomData;
9
10use crate::sql::{SQL, Token};
11use crate::traits::SQLParam;
12use crate::types::{BooleanLike, Compatible, DataType};
13
14use super::{AggregateKind, Expr, Null, Nullability, SQLExpr};
15use crate::scope::ScopeOnly;
16
17// =============================================================================
18// Entry Point
19// =============================================================================
20
21/// Starts a searched `CASE` expression.
22///
23/// Add at least one branch with [`when`](CaseInit::when), then finish with
24/// [`r#else`](CaseBuilder#method.else) or [`end`](CaseBuilder::end). The first
25/// branch's result sets the type of the whole expression; later results must
26/// have a compatible type.
27///
28/// # Examples
29///
30/// ```rust
31/// # use drizzle_core::dialect::{Dialect, DialectTypes, SQLiteDialect as D};
32/// # use drizzle_core::{ColumnRef, SQL, SQLParam, expr::*};
33/// # #[derive(Clone, Debug)] struct Value(String);
34/// # impl SQLParam for Value { const DIALECT: Dialect = Dialect::SQLite; type DialectMarker = D; }
35/// # impl<X: ToString> From<X> for Value { fn from(v: X) -> Self { Value(v.to_string()) } }
36/// # impl From<Value> for std::borrow::Cow<'_, Value> { fn from(v: Value) -> Self { Self::Owned(v) } }
37/// # type C<X, N = NonNull> = &'static SQLExpr<'static, Value, X, N>;
38/// # fn col<X: drizzle_core::types::DataType, N: Nullability>(c: &'static str) -> C<X, N> { Box::leak(Box::new(SQLExpr::new(SQL::column(ColumnRef::sql("users", c))))) }
39/// # type Int = <D as DialectTypes>::Int; type Text = <D as DialectTypes>::Text; type Real = <D as DialectTypes>::Double;
40/// # struct Users { id: C<Int>, age: C<Int>, name: C<Text>, email: C<Text, Null>, score: C<Real, Null>, active: C<<D as DialectTypes>::Bool>, created_at: C<<D as DialectTypes>::Timestamp> }
41/// # let users = Users { id: col("id"), age: col("age"), name: col("name"), email: col("email"), score: col("score"), active: col("active"), created_at: col("created_at") };
42/// let group = case()
43/// .when(gt(users.age, 65), "senior")
44/// .when(gt(users.age, 17), "adult")
45/// .r#else("minor");
46/// assert_eq!(
47/// group.sql(),
48/// r#"CASE WHEN "users"."age" > ? THEN ? WHEN "users"."age" > ? THEN ? ELSE ? END"#
49/// );
50/// ```
51///
52/// # Type safety
53///
54/// Branches with different result types do not compile:
55///
56/// ```rust,compile_fail
57/// # use drizzle_core::dialect::{Dialect, DialectTypes, SQLiteDialect as D};
58/// # use drizzle_core::{ColumnRef, SQL, SQLParam, expr::*};
59/// # #[derive(Clone, Debug)] struct Value(String);
60/// # impl SQLParam for Value { const DIALECT: Dialect = Dialect::SQLite; type DialectMarker = D; }
61/// # impl<X: ToString> From<X> for Value { fn from(v: X) -> Self { Value(v.to_string()) } }
62/// # impl From<Value> for std::borrow::Cow<'_, Value> { fn from(v: Value) -> Self { Self::Owned(v) } }
63/// # type C<X, N = NonNull> = &'static SQLExpr<'static, Value, X, N>;
64/// # fn col<X: drizzle_core::types::DataType, N: Nullability>(c: &'static str) -> C<X, N> { Box::leak(Box::new(SQLExpr::new(SQL::column(ColumnRef::sql("users", c))))) }
65/// # type Int = <D as DialectTypes>::Int; type Text = <D as DialectTypes>::Text; type Real = <D as DialectTypes>::Double;
66/// # struct Users { id: C<Int>, age: C<Int>, name: C<Text>, email: C<Text, Null>, score: C<Real, Null>, active: C<<D as DialectTypes>::Bool>, created_at: C<<D as DialectTypes>::Timestamp> }
67/// # let users = Users { id: col("id"), age: col("age"), name: col("name"), email: col("email"), score: col("score"), active: col("active"), created_at: col("created_at") };
68/// let wrong = case().when(gt(users.age, 65), "senior").r#else(0);
69/// ```
70#[must_use]
71pub fn case<'a, V: SQLParam>() -> CaseInit<'a, V> {
72 CaseInit {
73 sql: SQL::from(Token::CASE),
74 _marker: PhantomData,
75 }
76}
77
78// =============================================================================
79// CaseInit — before the first WHEN (no type established yet)
80// =============================================================================
81
82/// A `CASE` builder before its first branch; created by [`case`].
83///
84/// The result type is set by the first [`when`](Self::when).
85pub struct CaseInit<'a, V: SQLParam> {
86 sql: SQL<'a, V>,
87 _marker: PhantomData<V>,
88}
89
90impl<'a, V: SQLParam + 'a> CaseInit<'a, V> {
91 /// Adds the first `WHEN condition THEN result` branch.
92 ///
93 /// `condition` must be boolean. `result` sets the type of the whole `CASE`.
94 #[allow(clippy::type_complexity)]
95 pub fn when<C, R>(
96 self,
97 condition: C,
98 result: R,
99 ) -> CaseBuilder<
100 'a,
101 V,
102 R::SQLType,
103 R::Nullable,
104 <C::Aggregate as AggregateKind>::Or<R::Aggregate>,
105 (ScopeOnly<C::Sources>, R::Sources),
106 >
107 where
108 C: Expr<'a, V>,
109 R: Expr<'a, V>,
110 C::SQLType: BooleanLike,
111 {
112 let sql = self
113 .sql
114 .push(Token::WHEN)
115 .append(condition.into_expr_sql())
116 .push(Token::THEN)
117 .append(result.into_expr_sql());
118
119 CaseBuilder {
120 sql,
121 _marker: PhantomData,
122 }
123 }
124}
125
126// =============================================================================
127// CaseBuilder — after at least one WHEN (type T established)
128// =============================================================================
129
130/// A `CASE` builder with at least one branch.
131///
132/// `T` is the result type set by the first branch and `N` the nullability of
133/// the results so far. `S` records the tables read so far: `WHEN` conditions
134/// are only scope-checked (a NULL condition falls through to the next
135/// branch), while `THEN` results can make the result NULL.
136pub struct CaseBuilder<'a, V: SQLParam, T: DataType, N: Nullability, A: AggregateKind, S = ()> {
137 sql: SQL<'a, V>,
138 _marker: super::TypeMarker<(V, T, N, A, S)>,
139}
140
141impl<'a, V, T, N, A, S> CaseBuilder<'a, V, T, N, A, S>
142where
143 V: SQLParam + 'a,
144 T: DataType,
145 N: Nullability,
146 A: AggregateKind,
147{
148 /// Adds another `WHEN condition THEN result` branch.
149 ///
150 /// `condition` must be boolean and `result` must have a type compatible with
151 /// the first branch. The result becomes nullable if this branch's result is.
152 #[allow(clippy::type_complexity)]
153 pub fn when<C, R>(
154 self,
155 condition: C,
156 result: R,
157 ) -> CaseBuilder<
158 'a,
159 V,
160 T,
161 <N as Nullability>::Or<R::Nullable>,
162 <<A as AggregateKind>::Or<C::Aggregate> as AggregateKind>::Or<R::Aggregate>,
163 (S, (ScopeOnly<C::Sources>, R::Sources)),
164 >
165 where
166 C: Expr<'a, V>,
167 R: Expr<'a, V>,
168 C::SQLType: BooleanLike,
169 T: Compatible<R::SQLType>,
170 N: Nullability,
171 R::Nullable: Nullability,
172 A: AggregateKind,
173 C::Aggregate: AggregateKind,
174 R::Aggregate: AggregateKind,
175 {
176 let sql = self
177 .sql
178 .push(Token::WHEN)
179 .append(condition.into_expr_sql())
180 .push(Token::THEN)
181 .append(result.into_expr_sql());
182
183 CaseBuilder {
184 sql,
185 _marker: PhantomData,
186 }
187 }
188
189 /// Finishes the expression without `ELSE` (`... END`).
190 ///
191 /// Rows that match no branch give NULL, so the result is always nullable.
192 ///
193 /// # Examples
194 ///
195 /// ```rust
196 /// # use drizzle_core::dialect::{Dialect, DialectTypes, SQLiteDialect as D};
197 /// # use drizzle_core::{ColumnRef, SQL, SQLParam, expr::*};
198 /// # #[derive(Clone, Debug)] struct Value(String);
199 /// # impl SQLParam for Value { const DIALECT: Dialect = Dialect::SQLite; type DialectMarker = D; }
200 /// # impl<X: ToString> From<X> for Value { fn from(v: X) -> Self { Value(v.to_string()) } }
201 /// # impl From<Value> for std::borrow::Cow<'_, Value> { fn from(v: Value) -> Self { Self::Owned(v) } }
202 /// # type C<X, N = NonNull> = &'static SQLExpr<'static, Value, X, N>;
203 /// # fn col<X: drizzle_core::types::DataType, N: Nullability>(c: &'static str) -> C<X, N> { Box::leak(Box::new(SQLExpr::new(SQL::column(ColumnRef::sql("users", c))))) }
204 /// # type Int = <D as DialectTypes>::Int; type Text = <D as DialectTypes>::Text; type Real = <D as DialectTypes>::Double;
205 /// # struct Users { id: C<Int>, age: C<Int>, name: C<Text>, email: C<Text, Null>, score: C<Real, Null>, active: C<<D as DialectTypes>::Bool>, created_at: C<<D as DialectTypes>::Timestamp> }
206 /// # let users = Users { id: col("id"), age: col("age"), name: col("name"), email: col("email"), score: col("score"), active: col("active"), created_at: col("created_at") };
207 /// let label = case().when(users.active, "active").end();
208 /// assert_eq!(label.sql(), r#"CASE WHEN "users"."active" THEN ? END"#);
209 /// ```
210 pub fn end(self) -> SQLExpr<'a, V, T, Null, A, S> {
211 let sql = self.sql.push(Token::END);
212 SQLExpr::new(sql)
213 }
214
215 /// Finishes the expression with a default (`... ELSE default END`).
216 ///
217 /// `default` must have a type compatible with the first branch. The result is
218 /// nullable if any branch result or the default is.
219 ///
220 /// # Examples
221 ///
222 /// ```rust
223 /// # use drizzle_core::dialect::{Dialect, DialectTypes, SQLiteDialect as D};
224 /// # use drizzle_core::{ColumnRef, SQL, SQLParam, expr::*};
225 /// # #[derive(Clone, Debug)] struct Value(String);
226 /// # impl SQLParam for Value { const DIALECT: Dialect = Dialect::SQLite; type DialectMarker = D; }
227 /// # impl<X: ToString> From<X> for Value { fn from(v: X) -> Self { Value(v.to_string()) } }
228 /// # impl From<Value> for std::borrow::Cow<'_, Value> { fn from(v: Value) -> Self { Self::Owned(v) } }
229 /// # type C<X, N = NonNull> = &'static SQLExpr<'static, Value, X, N>;
230 /// # fn col<X: drizzle_core::types::DataType, N: Nullability>(c: &'static str) -> C<X, N> { Box::leak(Box::new(SQLExpr::new(SQL::column(ColumnRef::sql("users", c))))) }
231 /// # type Int = <D as DialectTypes>::Int; type Text = <D as DialectTypes>::Text; type Real = <D as DialectTypes>::Double;
232 /// # struct Users { id: C<Int>, age: C<Int>, name: C<Text>, email: C<Text, Null>, score: C<Real, Null>, active: C<<D as DialectTypes>::Bool>, created_at: C<<D as DialectTypes>::Timestamp> }
233 /// # let users = Users { id: col("id"), age: col("age"), name: col("name"), email: col("email"), score: col("score"), active: col("active"), created_at: col("created_at") };
234 /// let label = case().when(users.active, "active").r#else("inactive");
235 /// assert_eq!(label.sql(), r#"CASE WHEN "users"."active" THEN ? ELSE ? END"#);
236 /// ```
237 #[allow(clippy::type_complexity)]
238 pub fn r#else<D>(
239 self,
240 default: D,
241 ) -> SQLExpr<
242 'a,
243 V,
244 T,
245 <N as Nullability>::Or<D::Nullable>,
246 <A as AggregateKind>::Or<D::Aggregate>,
247 (S, D::Sources),
248 >
249 where
250 D: Expr<'a, V>,
251 T: Compatible<D::SQLType>,
252 N: Nullability,
253 D::Nullable: Nullability,
254 A: AggregateKind,
255 D::Aggregate: AggregateKind,
256 {
257 let sql = self
258 .sql
259 .push(Token::ELSE)
260 .append(default.into_expr_sql())
261 .push(Token::END);
262 SQLExpr::new(sql)
263 }
264}