use crate::dialect::{Dialect, MySqlDialect};
use crate::Value;
use std::marker::PhantomData;
pub trait Column<M>: Send + Sync + Clone {
fn name(&self) -> &'static str;
fn table(&self) -> &'static str;
}
#[derive(Debug, Clone)]
pub enum WhereClause {
Eq(String, Value),
Ne(String, Value),
Gt(String, Value),
Ge(String, Value),
Lt(String, Value),
Le(String, Value),
Like(String, Value),
IsNull(String),
IsNotNull(String),
In(String, Vec<Value>),
NotIn(String, Vec<Value>),
Between(String, Value, Value),
Raw(String),
}
impl WhereClause {
fn render(&self, dialect: &dyn Dialect) -> String {
match self {
WhereClause::Eq(col, v) => format!(
"{} = {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::Ne(col, v) => format!(
"{} != {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::Gt(col, v) => format!(
"{} > {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::Ge(col, v) => format!(
"{} >= {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::Lt(col, v) => format!(
"{} < {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::Le(col, v) => format!(
"{} <= {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::Like(col, v) => format!(
"{} LIKE {}",
dialect.quote(col),
v.to_param_with_dialect(dialect)
),
WhereClause::IsNull(col) => format!("{} IS NULL", dialect.quote(col)),
WhereClause::IsNotNull(col) => format!("{} IS NOT NULL", dialect.quote(col)),
WhereClause::In(col, vs) => {
let values: Vec<String> = vs
.iter()
.map(|v| v.to_param_with_dialect(dialect).to_string())
.collect();
format!("{} IN ({})", dialect.quote(col), values.join(", "))
}
WhereClause::NotIn(col, vs) => {
let values: Vec<String> = vs
.iter()
.map(|v| v.to_param_with_dialect(dialect).to_string())
.collect();
format!("{} NOT IN ({})", dialect.quote(col), values.join(", "))
}
WhereClause::Between(col, a, b) => format!(
"{} BETWEEN {} AND {}",
dialect.quote(col),
a.to_param_with_dialect(dialect),
b.to_param_with_dialect(dialect)
),
WhereClause::Raw(sql) => sql.clone(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OrderDirection {
Asc,
Desc,
}
#[derive(Debug, Clone)]
pub struct OrderBy {
pub column: String,
pub direction: OrderDirection,
}
pub struct LambdaWrapper<M> {
table: String,
selects: Vec<String>,
wheres: Vec<WhereClause>,
orders: Vec<OrderBy>,
limit: Option<u64>,
offset: Option<u64>,
dialect: Box<dyn Dialect>,
_marker: PhantomData<M>,
}
impl<M> LambdaWrapper<M> {
pub fn new(table: impl Into<String>) -> Self {
Self {
table: table.into(),
selects: Vec::new(),
wheres: Vec::new(),
orders: Vec::new(),
limit: None,
offset: None,
dialect: Box::new(MySqlDialect),
_marker: PhantomData,
}
}
pub fn with_dialect(table: impl Into<String>, dialect: Box<dyn Dialect>) -> Self {
Self {
table: table.into(),
selects: Vec::new(),
wheres: Vec::new(),
orders: Vec::new(),
limit: None,
offset: None,
dialect,
_marker: PhantomData,
}
}
pub fn select<C: Column<M>>(&mut self, col: C) -> &mut Self {
let name = col.name();
if crate::sql_safety::validate_identifier(name, "lambda select column").is_ok() {
self.selects.push(name.to_string());
}
self
}
pub fn select_many<C: Column<M>>(&mut self, cols: &[C]) -> &mut Self {
for c in cols {
let name = c.name();
if crate::sql_safety::validate_identifier(name, "lambda select column").is_ok() {
self.selects.push(name.to_string());
}
}
self
}
pub fn select_all(&mut self) -> &mut Self {
self.selects.clear();
self
}
pub fn eq<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Eq(col.name().to_string(), value));
self
}
pub fn ne<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Ne(col.name().to_string(), value));
self
}
pub fn gt<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Gt(col.name().to_string(), value));
self
}
pub fn ge<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Ge(col.name().to_string(), value));
self
}
pub fn lt<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Lt(col.name().to_string(), value));
self
}
pub fn le<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Le(col.name().to_string(), value));
self
}
pub fn like<C: Column<M>>(&mut self, col: C, value: Value) -> &mut Self {
self.wheres
.push(WhereClause::Like(col.name().to_string(), value));
self
}
pub fn is_null<C: Column<M>>(&mut self, col: C) -> &mut Self {
self.wheres
.push(WhereClause::IsNull(col.name().to_string()));
self
}
pub fn is_not_null<C: Column<M>>(&mut self, col: C) -> &mut Self {
self.wheres
.push(WhereClause::IsNotNull(col.name().to_string()));
self
}
pub fn r#in<C: Column<M>>(&mut self, col: C, values: Vec<Value>) -> &mut Self {
self.wheres
.push(WhereClause::In(col.name().to_string(), values));
self
}
pub fn not_in<C: Column<M>>(&mut self, col: C, values: Vec<Value>) -> &mut Self {
self.wheres
.push(WhereClause::NotIn(col.name().to_string(), values));
self
}
pub fn between<C: Column<M>>(&mut self, col: C, a: Value, b: Value) -> &mut Self {
self.wheres
.push(WhereClause::Between(col.name().to_string(), a, b));
self
}
pub fn raw_where(&mut self, sql: impl Into<String>) -> &mut Self {
self.wheres.push(WhereClause::Raw(sql.into()));
self
}
pub fn order_by_asc<C: Column<M>>(&mut self, col: C) -> &mut Self {
self.orders.push(OrderBy {
column: col.name().to_string(),
direction: OrderDirection::Asc,
});
self
}
pub fn order_by_desc<C: Column<M>>(&mut self, col: C) -> &mut Self {
self.orders.push(OrderBy {
column: col.name().to_string(),
direction: OrderDirection::Desc,
});
self
}
pub fn limit(&mut self, n: u64) -> &mut Self {
self.limit = Some(n);
self
}
pub fn offset(&mut self, n: u64) -> &mut Self {
self.offset = Some(n);
self
}
pub fn page(&mut self, page: u64, page_size: u64) -> &mut Self {
self.limit = Some(page_size);
if page > 1 {
self.offset = Some((page - 1) * page_size);
} else {
self.offset = None;
}
self
}
pub fn build_select(&self) -> String {
let quoted_table = self.dialect.quote(&self.table);
let select_sql = if self.selects.is_empty() {
"*".to_string()
} else {
self.selects
.iter()
.map(|c| self.dialect.quote(c))
.collect::<Vec<_>>()
.join(", ")
};
let mut sql = format!("SELECT {} FROM {}", select_sql, quoted_table);
if !self.wheres.is_empty() {
let conditions: Vec<String> = self
.wheres
.iter()
.map(|w| w.render(self.dialect.as_ref()))
.collect();
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
if !self.orders.is_empty() {
let orders: Vec<String> = self
.orders
.iter()
.map(|o| {
let dir = match o.direction {
OrderDirection::Asc => "ASC",
OrderDirection::Desc => "DESC",
};
format!("{} {}", self.dialect.quote(&o.column), dir)
})
.collect();
sql.push_str(" ORDER BY ");
sql.push_str(&orders.join(", "));
}
if let Some(l) = self.limit {
sql.push_str(&format!(" LIMIT {}", l));
}
if let Some(o) = self.offset {
sql.push_str(&format!(" OFFSET {}", o));
}
sql
}
pub fn build_count(&self) -> String {
let quoted_table = self.dialect.quote(&self.table);
let mut sql = format!("SELECT COUNT(*) FROM {}", quoted_table);
if !self.wheres.is_empty() {
let conditions: Vec<String> = self
.wheres
.iter()
.map(|w| w.render(self.dialect.as_ref()))
.collect();
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
sql
}
pub fn build_exists(&self) -> String {
let inner = self.build_select();
let inner = if let Some(pos) = inner.find(" FROM ") {
format!("SELECT 1{}", &inner[pos..])
} else {
inner
};
format!("SELECT EXISTS({}) AS exists_flag", inner)
}
pub fn build_delete(&self) -> String {
let quoted_table = self.dialect.quote(&self.table);
let mut sql = format!("DELETE FROM {}", quoted_table);
if !self.wheres.is_empty() {
let conditions: Vec<String> = self
.wheres
.iter()
.map(|w| w.render(self.dialect.as_ref()))
.collect();
sql.push_str(" WHERE ");
sql.push_str(&conditions.join(" AND "));
}
sql
}
pub fn where_count(&self) -> usize {
self.wheres.len()
}
pub fn select_count(&self) -> usize {
self.selects.len()
}
pub fn table(&self) -> &str {
&self.table
}
pub fn reset(&mut self) -> &mut Self {
self.selects.clear();
self.wheres.clear();
self.orders.clear();
self.limit = None;
self.offset = None;
self
}
}
#[macro_export]
macro_rules! define_columns {
(
$columns_struct:ident for $model:ident table = $table:literal {
$( $field:ident => $name:literal ),* $(,)?
}
) => {
#[derive(Debug, Clone, Copy)]
pub struct $columns_struct {
/// 字段名
pub name: &'static str,
pub table: &'static str,
}
impl $crate::lambda::Column<$model> for $columns_struct {
fn name(&self) -> &'static str {
self.name
}
fn table(&self) -> &'static str {
self.table
}
}
impl $columns_struct {
$(
#[allow(non_upper_case_globals, dead_code)]
pub const $field: $columns_struct = $columns_struct { name: $name, table: $table };
)*
}
};
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dialect::PostgreSqlDialect;
use crate::get_dialect;
use crate::DbType;
struct User;
define_columns! {
UserColumns for User table = "users" {
Id => "id",
Name => "name",
Age => "age",
Email => "email",
}
}
struct Order;
define_columns! {
OrderColumns for Order table = "orders" {
OrderId => "order_id",
UserId => "user_id",
Total => "total",
}
}
#[test]
fn test_column_name_and_table() {
assert_eq!(UserColumns::Id.name(), "id");
assert_eq!(UserColumns::Id.table(), "users");
assert_eq!(UserColumns::Name.name(), "name");
assert_eq!(UserColumns::Age.name(), "age");
assert_eq!(UserColumns::Email.name(), "email");
}
#[test]
fn test_column_for_different_models() {
assert_eq!(OrderColumns::OrderId.name(), "order_id");
assert_eq!(OrderColumns::OrderId.table(), "orders");
assert_eq!(OrderColumns::UserId.name(), "user_id");
}
#[test]
fn test_new_wrapper() {
let w = LambdaWrapper::<User>::new("users");
assert_eq!(w.table(), "users");
assert_eq!(w.where_count(), 0);
assert_eq!(w.select_count(), 0);
}
#[test]
fn test_select_single() {
let mut w = LambdaWrapper::<User>::new("users");
w.select(UserColumns::Id);
let sql = w.build_select();
assert!(sql.contains("SELECT `id` FROM `users`"));
}
#[test]
fn test_select_multiple() {
let mut w = LambdaWrapper::<User>::new("users");
w.select(UserColumns::Id)
.select(UserColumns::Name)
.select(UserColumns::Age);
let sql = w.build_select();
assert!(sql.contains("`id`, `name`, `age`"));
}
#[test]
fn test_select_many() {
let mut w = LambdaWrapper::<User>::new("users");
w.select_many(&[UserColumns::Id, UserColumns::Name, UserColumns::Age]);
let sql = w.build_select();
assert!(sql.contains("`id`, `name`, `age`"));
}
#[test]
fn test_select_all_clears_selects() {
let mut w = LambdaWrapper::<User>::new("users");
w.select(UserColumns::Id);
assert_eq!(w.select_count(), 1);
w.select_all();
assert_eq!(w.select_count(), 0);
let sql = w.build_select();
assert!(sql.contains("SELECT * FROM"));
}
#[test]
fn test_default_select_is_star() {
let w = LambdaWrapper::<User>::new("users");
let sql = w.build_select();
assert!(sql.contains("SELECT * FROM `users`"));
}
#[test]
fn test_where_eq() {
let mut w = LambdaWrapper::<User>::new("users");
w.eq(UserColumns::Id, Value::I64(1));
let sql = w.build_select();
assert!(sql.contains("WHERE `id` = 1"));
}
#[test]
fn test_where_ne() {
let mut w = LambdaWrapper::<User>::new("users");
w.ne(UserColumns::Id, Value::I64(1));
let sql = w.build_select();
assert!(sql.contains("`id` != 1"));
}
#[test]
fn test_where_gt_ge_lt_le() {
let mut w = LambdaWrapper::<User>::new("users");
w.gt(UserColumns::Age, Value::I64(18))
.ge(UserColumns::Age, Value::I64(20))
.lt(UserColumns::Age, Value::I64(65))
.le(UserColumns::Age, Value::I64(60));
let sql = w.build_select();
assert!(sql.contains("`age` > 18"));
assert!(sql.contains("`age` >= 20"));
assert!(sql.contains("`age` < 65"));
assert!(sql.contains("`age` <= 60"));
}
#[test]
fn test_where_like() {
let mut w = LambdaWrapper::<User>::new("users");
w.like(UserColumns::Name, Value::String("%alice%".to_string()));
let sql = w.build_select();
assert!(sql.contains("`name` LIKE '%alice%'"));
}
#[test]
fn test_where_is_null() {
let mut w = LambdaWrapper::<User>::new("users");
w.is_null(UserColumns::Email);
let sql = w.build_select();
assert!(sql.contains("`email` IS NULL"));
}
#[test]
fn test_where_is_not_null() {
let mut w = LambdaWrapper::<User>::new("users");
w.is_not_null(UserColumns::Email);
let sql = w.build_select();
assert!(sql.contains("`email` IS NOT NULL"));
}
#[test]
fn test_where_in() {
let mut w = LambdaWrapper::<User>::new("users");
w.r#in(
UserColumns::Id,
vec![Value::I64(1), Value::I64(2), Value::I64(3)],
);
let sql = w.build_select();
assert!(sql.contains("`id` IN (1, 2, 3)"));
}
#[test]
fn test_where_not_in() {
let mut w = LambdaWrapper::<User>::new("users");
w.not_in(UserColumns::Id, vec![Value::I64(1), Value::I64(2)]);
let sql = w.build_select();
assert!(sql.contains("`id` NOT IN (1, 2)"));
}
#[test]
fn test_where_between() {
let mut w = LambdaWrapper::<User>::new("users");
w.between(UserColumns::Age, Value::I64(18), Value::I64(65));
let sql = w.build_select();
assert!(sql.contains("`age` BETWEEN 18 AND 65"));
}
#[test]
fn test_where_multiple_anded() {
let mut w = LambdaWrapper::<User>::new("users");
w.eq(UserColumns::Id, Value::I64(1))
.gt(UserColumns::Age, Value::I64(18))
.like(UserColumns::Name, Value::String("alice%".to_string()));
let sql = w.build_select();
assert!(sql.contains("`id` = 1"));
assert!(sql.contains("`age` > 18"));
assert!(sql.contains("`name` LIKE 'alice%'"));
assert!(sql.contains(" AND "));
}
#[test]
fn test_where_raw() {
let mut w = LambdaWrapper::<User>::new("users");
w.raw_where("name = 'alice' OR name = 'bob'");
let sql = w.build_select();
assert!(sql.contains("name = 'alice' OR name = 'bob'"));
}
#[test]
fn test_order_by_asc() {
let mut w = LambdaWrapper::<User>::new("users");
w.order_by_asc(UserColumns::Name);
let sql = w.build_select();
assert!(sql.contains("ORDER BY `name` ASC"));
}
#[test]
fn test_order_by_desc() {
let mut w = LambdaWrapper::<User>::new("users");
w.order_by_desc(UserColumns::Id);
let sql = w.build_select();
assert!(sql.contains("ORDER BY `id` DESC"));
}
#[test]
fn test_order_by_multiple() {
let mut w = LambdaWrapper::<User>::new("users");
w.order_by_asc(UserColumns::Name)
.order_by_desc(UserColumns::Id);
let sql = w.build_select();
assert!(sql.contains("ORDER BY `name` ASC, `id` DESC"));
}
#[test]
fn test_limit() {
let mut w = LambdaWrapper::<User>::new("users");
w.limit(10);
let sql = w.build_select();
assert!(sql.contains("LIMIT 10"));
}
#[test]
fn test_offset() {
let mut w = LambdaWrapper::<User>::new("users");
w.limit(10).offset(20);
let sql = w.build_select();
assert!(sql.contains("LIMIT 10"));
assert!(sql.contains("OFFSET 20"));
}
#[test]
fn test_page() {
let mut w = LambdaWrapper::<User>::new("users");
w.page(3, 20); let sql = w.build_select();
assert!(sql.contains("LIMIT 20"));
assert!(sql.contains("OFFSET 40")); }
#[test]
fn test_page_1_no_offset() {
let mut w = LambdaWrapper::<User>::new("users");
w.page(1, 10);
let sql = w.build_select();
assert!(sql.contains("LIMIT 10"));
assert!(!sql.contains("OFFSET")); }
#[test]
fn test_build_count() {
let mut w = LambdaWrapper::<User>::new("users");
w.gt(UserColumns::Age, Value::I64(18));
let sql = w.build_count();
assert!(sql.contains("SELECT COUNT(*) FROM `users`"));
assert!(sql.contains("`age` > 18"));
assert!(!sql.contains("ORDER BY"));
assert!(!sql.contains("LIMIT"));
}
#[test]
fn test_build_exists() {
let mut w = LambdaWrapper::<User>::new("users");
w.eq(UserColumns::Id, Value::I64(1));
let sql = w.build_exists();
assert!(sql.starts_with("SELECT EXISTS("));
assert!(sql.contains("SELECT 1 FROM `users`"));
assert!(sql.contains("`id` = 1"));
assert!(sql.ends_with(") AS exists_flag"));
}
#[test]
fn test_build_delete() {
let mut w = LambdaWrapper::<User>::new("users");
w.eq(UserColumns::Id, Value::I64(1));
let sql = w.build_delete();
assert!(sql.starts_with("DELETE FROM `users`"));
assert!(sql.contains("WHERE `id` = 1"));
}
#[test]
fn test_with_postgres_dialect() {
let dialect: Box<dyn Dialect> = Box::new(PostgreSqlDialect);
let mut w = LambdaWrapper::<User>::with_dialect("users", dialect);
w.eq(UserColumns::Id, Value::I64(1));
let sql = w.build_select();
assert!(sql.contains("\"users\""));
assert!(sql.contains("\"id\" = 1"));
}
#[test]
fn test_postgres_select() {
let dialect: Box<dyn Dialect> = Box::new(PostgreSqlDialect);
let mut w = LambdaWrapper::<User>::with_dialect("users", dialect);
w.select(UserColumns::Id).select(UserColumns::Name);
let sql = w.build_select();
assert!(sql.contains("\"id\", \"name\""));
}
#[test]
fn test_complex_query() {
let mut w = LambdaWrapper::<User>::new("users");
w.select(UserColumns::Id)
.select(UserColumns::Name)
.select(UserColumns::Age)
.gt(UserColumns::Age, Value::I64(18))
.like(UserColumns::Name, Value::String("a%".to_string()))
.is_not_null(UserColumns::Email)
.order_by_desc(UserColumns::Id)
.limit(10)
.offset(20);
let sql = w.build_select();
assert!(sql.contains("SELECT `id`, `name`, `age` FROM `users`"));
assert!(sql.contains("`age` > 18"));
assert!(sql.contains("`name` LIKE 'a%'"));
assert!(sql.contains("`email` IS NOT NULL"));
assert!(sql.contains("ORDER BY `id` DESC"));
assert!(sql.contains("LIMIT 10"));
assert!(sql.contains("OFFSET 20"));
}
#[test]
fn test_reset_clears_all() {
let mut w = LambdaWrapper::<User>::new("users");
w.select(UserColumns::Id)
.eq(UserColumns::Id, Value::I64(1))
.order_by_asc(UserColumns::Name)
.limit(10);
w.reset();
assert_eq!(w.select_count(), 0);
assert_eq!(w.where_count(), 0);
let sql = w.build_select();
assert!(sql.contains("SELECT * FROM `users`"));
assert!(!sql.contains("WHERE"));
assert!(!sql.contains("ORDER BY"));
assert!(!sql.contains("LIMIT"));
}
#[test]
fn test_different_models_dont_share_columns() {
let mut user_w = LambdaWrapper::<User>::new("users");
user_w.eq(UserColumns::Id, Value::I64(1));
let mut order_w = LambdaWrapper::<Order>::new("orders");
order_w.eq(OrderColumns::OrderId, Value::I64(100));
let user_sql = user_w.build_select();
let order_sql = order_w.build_select();
assert!(user_sql.contains("`users`"));
assert!(user_sql.contains("`id` = 1"));
assert!(order_sql.contains("`orders`"));
assert!(order_sql.contains("`order_id` = 100"));
}
#[test]
fn test_type_safety_compile_time_check() {
let mut w = LambdaWrapper::<User>::new("users");
w.eq(UserColumns::Id, Value::I64(1));
assert_eq!(w.where_count(), 1);
}
#[test]
fn test_string_value_escape() {
let mut w = LambdaWrapper::<User>::new("users");
w.eq(UserColumns::Name, Value::String("O'Brien".to_string()));
let sql = w.build_select();
assert!(sql.contains("'O\\'Brien'"));
let pg_dialect: Box<dyn Dialect> = get_dialect(DbType::PostgreSQL).unwrap();
let mut w_pg = LambdaWrapper::<User>::with_dialect("users", pg_dialect);
w_pg.eq(UserColumns::Name, Value::String("O'Brien".to_string()));
let sql_pg = w_pg.build_select();
assert!(sql_pg.contains("'O''Brien'"));
}
#[test]
fn test_with_real_mysql_dialect() {
let dialect: Box<dyn Dialect> = get_dialect(DbType::MySQL).unwrap();
let mut w = LambdaWrapper::<User>::with_dialect("users", dialect);
w.select(UserColumns::Id)
.select(UserColumns::Name)
.eq(UserColumns::Id, Value::I64(42))
.order_by_desc(UserColumns::Id);
let sql = w.build_select();
assert!(sql.contains("SELECT `id`, `name` FROM `users`"));
assert!(sql.contains("`id` = 42"));
assert!(sql.contains("ORDER BY `id` DESC"));
}
}