use crate::expr::Expr;
use crate::expr::node::ExprNode;
use std::marker::PhantomData;
pub struct Case<V> {
_never: PhantomData<fn() -> V>,
}
impl<V> Case<V> {
pub fn when(cond: Expr<bool>, then_val: Expr<V>) -> CaseBuilder<V> {
CaseBuilder {
arms: vec![(cond.node, then_val.node)],
_v: PhantomData,
}
}
}
#[must_use = "CaseBuilder must be closed with .otherwise(default) to produce an Expr<V>"]
pub struct CaseBuilder<V> {
arms: Vec<(ExprNode, ExprNode)>,
_v: PhantomData<fn() -> V>,
}
impl<V> CaseBuilder<V> {
pub fn when(mut self, cond: Expr<bool>, then_val: Expr<V>) -> Self {
self.arms.push((cond.node, then_val.node));
self
}
pub fn otherwise(self, default: Expr<V>) -> Expr<V> {
Expr::from_node(ExprNode::Case {
arms: self.arms,
otherwise: Box::new(default.node),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Expr;
use crate::descriptor::ModelDescriptor;
use crate::expr::sql::emit_expr;
use crate::pg::accumulator::SqlAccumulator;
use crate::query::field::FieldRef;
use crate::query::portable::SqlEmitContext;
struct Acct;
impl crate::model::__sealed::Sealed for Acct {}
#[allow(clippy::manual_async_fn)]
impl crate::model::Model for Acct {
type Pk = i64;
type Fields = ();
fn table_name() -> &'static str {
"accts"
}
fn pk_value(&self) -> &i64 {
unreachable!()
}
fn descriptor() -> &'static ModelDescriptor {
unreachable!()
}
fn get(
_ctx: &mut crate::context::DjogiContext,
_id: i64,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn create(
_ctx: &mut crate::context::DjogiContext,
_v: Self,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn save<'ctx>(
&'ctx mut self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
fn delete(
self,
_ctx: &mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send {
async { unreachable!() }
}
fn refresh_from_db<'ctx>(
&'ctx self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
}
#[test]
fn single_arm_case_emits_when_then_else() {
let f: FieldRef<Acct, i64> = FieldRef::new("balance");
let expr = Case::when(
f.as_expr().lt(Expr::literal(0i64)),
Expr::literal("overdrawn".to_string()),
)
.otherwise(Expr::literal("ok".to_string()));
let mut qb = SqlAccumulator::new("");
emit_expr(&mut qb, &expr.node, SqlEmitContext::root())
.expect("case expression should lower to SQL");
let sql = qb.sql();
assert_eq!(
sql.trim(),
"CASE WHEN balance < $1 THEN $2 ELSE $3 END",
"got: {sql}"
);
}
#[test]
fn multi_arm_case_emits_arms_in_order() {
let f: FieldRef<Acct, i64> = FieldRef::new("balance");
let g: FieldRef<Acct, i64> = FieldRef::new("balance");
let expr = Case::when(
f.as_expr().lt(Expr::literal(0i64)),
Expr::literal("overdrawn".to_string()),
)
.when(
g.as_expr().eq(Expr::literal(0i64)),
Expr::literal("zero".to_string()),
)
.otherwise(Expr::literal("ok".to_string()));
let mut qb = SqlAccumulator::new("");
emit_expr(&mut qb, &expr.node, SqlEmitContext::root())
.expect("case expression should lower to SQL");
let sql = qb.sql();
assert_eq!(
sql.trim(),
"CASE WHEN balance < $1 THEN $2 WHEN balance = $3 THEN $4 ELSE $5 END",
"got: {sql}"
);
}
}