Skip to main content

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}