use crate::Dialect;
use crate::model::Model;
use crate::Value;
use std::fmt;
pub struct QueryBuilder<M: Model> {
table: Option<String>,
select_columns: Vec<String>,
where_conditions: Vec<WhereCondition>,
order_by: Vec<OrderClause>,
group_by: Vec<String>,
having_conditions: Vec<WhereCondition>,
limit_value: Option<usize>,
offset_value: Option<usize>,
joins: Vec<JoinClause>,
dialect: Box<dyn Dialect>,
soft_delete_disabled: bool,
tenant_id_value: Option<i64>,
tenant_disabled: bool,
keyset_cursor: Option<KeysetCursor>,
#[allow(dead_code)]
model: std::marker::PhantomData<M>,
}
#[derive(Debug, Clone)]
struct KeysetCursor {
field: String,
value: Value,
direction: KeysetDirection,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum KeysetDirection {
After,
Before,
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
enum WhereCondition {
And(String),
Or(String),
Eq(String, Value),
Ne(String, Value),
Gt(String, Value),
Ge(String, Value),
Lt(String, Value),
Le(String, Value),
Like(String, Value),
OrEq(String, Value),
OrNe(String, Value),
OrGt(String, Value),
OrGe(String, Value),
OrLt(String, Value),
OrLe(String, Value),
OrLike(String, Value),
In(String, Vec<Value>),
NotIn(String, Vec<Value>),
Between(String, Value, Value),
NotBetween(String, Value, Value),
Null(String),
NotNull(String),
Exists(String),
NotExists(String),
}
#[derive(Debug, Clone)]
struct OrderClause {
field: String,
direction: OrderDirection,
}
#[derive(Debug, Clone)]
enum OrderDirection {
Asc,
Desc,
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
enum JoinClause {
Inner(String, String, String),
Left(String, String, String),
Right(String, String, String),
Cross(String, String),
}
impl<M: Model> QueryBuilder<M> {
pub fn new(dialect: Box<dyn Dialect>) -> Self {
Self {
table: None,
select_columns: vec!["*".to_string()],
where_conditions: Vec::new(),
order_by: Vec::new(),
group_by: Vec::new(),
having_conditions: Vec::new(),
limit_value: None,
offset_value: None,
joins: Vec::new(),
dialect,
soft_delete_disabled: false,
tenant_id_value: None,
tenant_disabled: false,
keyset_cursor: None,
model: std::marker::PhantomData,
}
}
pub fn table(mut self, table: impl Into<String>) -> Self {
self.table = Some(table.into());
self
}
pub fn without_soft_delete(mut self) -> Self {
self.soft_delete_disabled = true;
self
}
pub fn is_soft_delete_disabled(&self) -> bool {
self.soft_delete_disabled
}
fn soft_delete_field(&self) -> Option<&'static str> {
if self.soft_delete_disabled {
return None;
}
M::soft_delete_field()
}
fn build_soft_delete_condition(&self) -> Option<String> {
self.soft_delete_field()
.map(|field| format!("{} IS NULL", self.dialect.quote(field)))
}
pub fn with_tenant_id(mut self, tenant_id: i64) -> Self {
self.tenant_id_value = Some(tenant_id);
self
}
pub fn without_tenant(mut self) -> Self {
self.tenant_disabled = true;
self
}
pub fn is_tenant_disabled(&self) -> bool {
self.tenant_disabled
}
fn tenant_field(&self) -> Option<&'static str> {
if self.tenant_disabled {
return None;
}
M::tenant_field()
}
fn tenant_id_value(&self) -> Option<i64> {
if self.tenant_disabled {
return None;
}
self.tenant_id_value
}
fn build_tenant_condition(&self) -> Option<(String, Value)> {
let field = self.tenant_field()?;
let tid = self.tenant_id_value()?;
Some((
format!("{} = ?", self.dialect.quote(field)),
Value::I64(tid),
))
}
pub fn select(mut self, columns: Vec<&str>) -> Self {
self.select_columns = columns.into_iter().map(|s| s.to_string()).collect();
self
}
pub fn select_quoted(mut self, columns: Vec<&str>) -> Result<Self, crate::DbError> {
let mut quoted = Vec::with_capacity(columns.len());
for col in columns {
crate::sql_safety::validate_identifier(col, "select column")?;
quoted.push(self.dialect.quote(col));
}
self.select_columns = quoted;
Ok(self)
}
#[deprecated(
since = "1.3.0",
note = "P0-2: 字符串拼接存在 SQL 注入风险,请使用 where_eq/where_ne/where_gt/where_lt/where_like 等参数化方法"
)]
pub fn where_cond(mut self, condition: impl Into<String>) -> Self {
self.where_conditions
.push(WhereCondition::And(condition.into()));
self
}
#[deprecated(
since = "1.3.0",
note = "P0-2: 字符串拼接存在 SQL 注入风险,请使用参数化方法"
)]
pub fn or_where(mut self, condition: impl Into<String>) -> Self {
self.where_conditions
.push(WhereCondition::Or(condition.into()));
self
}
pub fn where_eq(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::Eq(field.into(), value));
self
}
pub fn where_ne(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::Ne(field.into(), value));
self
}
pub fn where_gt(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::Gt(field.into(), value));
self
}
pub fn where_ge(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::Ge(field.into(), value));
self
}
pub fn where_lt(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::Lt(field.into(), value));
self
}
pub fn where_le(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::Le(field.into(), value));
self
}
pub fn where_like(mut self, field: impl Into<String>, pattern: Value) -> Self {
self.where_conditions
.push(WhereCondition::Like(field.into(), pattern));
self
}
pub fn or_where_eq(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrEq(field.into(), value));
self
}
pub fn or_where_ne(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrNe(field.into(), value));
self
}
pub fn or_where_gt(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrGt(field.into(), value));
self
}
pub fn or_where_ge(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrGe(field.into(), value));
self
}
pub fn or_where_lt(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrLt(field.into(), value));
self
}
pub fn or_where_le(mut self, field: impl Into<String>, value: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrLe(field.into(), value));
self
}
pub fn or_where_like(mut self, field: impl Into<String>, pattern: Value) -> Self {
self.where_conditions
.push(WhereCondition::OrLike(field.into(), pattern));
self
}
pub fn where_in(mut self, field: impl Into<String>, values: Vec<Value>) -> Self {
self.where_conditions
.push(WhereCondition::In(field.into(), values));
self
}
pub fn where_not_in(mut self, field: impl Into<String>, values: Vec<Value>) -> Self {
self.where_conditions
.push(WhereCondition::NotIn(field.into(), values));
self
}
pub fn where_between(mut self, field: impl Into<String>, start: Value, end: Value) -> Self {
self.where_conditions
.push(WhereCondition::Between(field.into(), start, end));
self
}
pub fn where_not_between(mut self, field: impl Into<String>, start: Value, end: Value) -> Self {
self.where_conditions
.push(WhereCondition::NotBetween(field.into(), start, end));
self
}
pub fn where_null(mut self, field: impl Into<String>) -> Self {
self.where_conditions
.push(WhereCondition::Null(field.into()));
self
}
pub fn where_not_null(mut self, field: impl Into<String>) -> Self {
self.where_conditions
.push(WhereCondition::NotNull(field.into()));
self
}
pub fn order_by(mut self, field: impl Into<String>) -> Self {
self.order_by.push(OrderClause {
field: field.into(),
direction: OrderDirection::Asc,
});
self
}
pub fn order_desc(mut self, field: impl Into<String>) -> Self {
self.order_by.push(OrderClause {
field: field.into(),
direction: OrderDirection::Desc,
});
self
}
pub fn group_by(mut self, field: impl Into<String>) -> Self {
self.group_by.push(field.into());
self
}
pub fn having(mut self, condition: impl Into<String>) -> Self {
self.having_conditions
.push(WhereCondition::And(condition.into()));
self
}
pub fn limit(mut self, limit: usize) -> Self {
self.limit_value = Some(limit);
self
}
pub fn offset(mut self, offset: usize) -> Self {
self.offset_value = Some(offset);
self
}
pub fn page(mut self, page: usize, page_size: usize) -> Self {
self.limit_value = Some(page_size);
self.offset_value = Some((page.saturating_sub(1)) * page_size);
self
}
pub fn keyset_after(
mut self,
field: impl Into<String>,
cursor_value: Value,
page_size: usize,
) -> Self {
let field_str = field.into();
if let Some(existing) = self.order_by.iter_mut().find(|o| o.field == field_str) {
existing.direction = OrderDirection::Asc;
} else {
self.order_by.push(OrderClause {
field: field_str.clone(),
direction: OrderDirection::Asc,
});
}
self.limit_value = Some(page_size);
self.offset_value = None;
self.keyset_cursor = Some(KeysetCursor {
field: field_str,
value: cursor_value,
direction: KeysetDirection::After,
});
self
}
pub fn keyset_before(
mut self,
field: impl Into<String>,
cursor_value: Value,
page_size: usize,
) -> Self {
let field_str = field.into();
if let Some(existing) = self.order_by.iter_mut().find(|o| o.field == field_str) {
existing.direction = OrderDirection::Desc;
} else {
self.order_by.push(OrderClause {
field: field_str.clone(),
direction: OrderDirection::Desc,
});
}
self.limit_value = Some(page_size);
self.offset_value = None;
self.keyset_cursor = Some(KeysetCursor {
field: field_str,
value: cursor_value,
direction: KeysetDirection::Before,
});
self
}
pub fn join_inner(
mut self,
table: impl Into<String>,
on_left: impl Into<String>,
on_right: impl Into<String>,
) -> Self {
self.joins.push(JoinClause::Inner(
table.into(),
on_left.into(),
on_right.into(),
));
self
}
pub fn join_left(
mut self,
table: impl Into<String>,
on_left: impl Into<String>,
on_right: impl Into<String>,
) -> Self {
self.joins.push(JoinClause::Left(
table.into(),
on_left.into(),
on_right.into(),
));
self
}
pub fn join_right(
mut self,
table: impl Into<String>,
on_left: impl Into<String>,
on_right: impl Into<String>,
) -> Self {
self.joins.push(JoinClause::Right(
table.into(),
on_left.into(),
on_right.into(),
));
self
}
#[tracing::instrument(skip(self), fields(op = "select"))]
pub fn build_select(&self) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let columns = if self.select_columns.is_empty() {
"*".to_string()
} else {
self.select_columns.join(", ")
};
let mut sql = format!("SELECT {} FROM {}", columns, self.dialect.quote(&table));
for join in &self.joins {
match join {
JoinClause::Inner(t, l, r) => {
sql.push_str(&format!(
" INNER JOIN {} ON {} = {}",
self.dialect.quote(t),
self.dialect.quote(l),
self.dialect.quote(r)
));
}
JoinClause::Left(t, l, r) => {
sql.push_str(&format!(
" LEFT JOIN {} ON {} = {}",
self.dialect.quote(t),
self.dialect.quote(l),
self.dialect.quote(r)
));
}
JoinClause::Right(t, l, r) => {
sql.push_str(&format!(
" RIGHT JOIN {} ON {} = {}",
self.dialect.quote(t),
self.dialect.quote(l),
self.dialect.quote(r)
));
}
JoinClause::Cross(t, on) => {
sql.push_str(&format!(
" CROSS JOIN {} ON {}",
self.dialect.quote(t),
self.dialect.quote(on)
));
}
}
}
let where_clause = self.build_where_clause();
if !where_clause.is_empty() {
sql.push_str(&where_clause);
}
if !self.group_by.is_empty() {
let cols: Vec<String> = self
.group_by
.iter()
.map(|c| self.dialect.quote(c))
.collect();
sql.push_str(" GROUP BY ");
sql.push_str(&cols.join(", "));
}
if !self.having_conditions.is_empty() {
sql.push_str(" HAVING ");
for (i, cond) in self.having_conditions.iter().enumerate() {
if i > 0 {
sql.push_str(" AND ");
}
if let WhereCondition::And(c) = cond {
sql.push_str(c);
}
}
}
if !self.order_by.is_empty() {
let order_cols: Vec<String> = self
.order_by
.iter()
.map(|o| {
let dir = match o.direction {
OrderDirection::Asc => " ASC",
OrderDirection::Desc => " DESC",
};
format!("{}{}", self.dialect.quote(&o.field), dir)
})
.collect();
sql.push_str(" ORDER BY ");
sql.push_str(&order_cols.join(", "));
}
if let Some(limit) = self.limit_value {
sql.push_str(&format!(" LIMIT {}", limit));
}
if let Some(offset) = self.offset_value {
sql.push_str(&format!(" OFFSET {}", offset));
}
sql
}
fn build_where_clause(&self) -> String {
self.build_where_clause_with_options(true)
}
fn build_where_clause_with_options(&self, include_soft_delete: bool) -> String {
let soft_delete_cond = if include_soft_delete {
self.build_soft_delete_condition()
} else {
None
};
let tenant_cond = self.build_tenant_condition().map(|(sql, value)| {
sql.replacen('?', &value.to_param_with_dialect(&*self.dialect), 1)
});
if self.where_conditions.is_empty()
&& soft_delete_cond.is_none()
&& tenant_cond.is_none()
&& self.keyset_cursor.is_none()
{
return String::new();
}
let mut conditions: Vec<String> = self
.where_conditions
.iter()
.map(|cond| match cond {
WhereCondition::And(c) => c.clone(),
WhereCondition::Or(c) => format!("OR {}", c),
WhereCondition::Eq(f, v) => format!(
"{} = {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::Ne(f, v) => format!(
"{} != {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::Gt(f, v) => format!(
"{} > {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::Ge(f, v) => format!(
"{} >= {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::Lt(f, v) => format!(
"{} < {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::Le(f, v) => format!(
"{} <= {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::Like(f, v) => format!(
"{} LIKE {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrEq(f, v) => format!(
"OR {} = {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrNe(f, v) => format!(
"OR {} != {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrGt(f, v) => format!(
"OR {} > {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrGe(f, v) => format!(
"OR {} >= {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrLt(f, v) => format!(
"OR {} < {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrLe(f, v) => format!(
"OR {} <= {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::OrLike(f, v) => format!(
"OR {} LIKE {}",
self.dialect.quote(f),
v.to_param_with_dialect(&*self.dialect)
),
WhereCondition::In(f, vals) => {
let vals_str: Vec<String> = vals
.iter()
.map(|v| v.to_param_with_dialect(&*self.dialect).to_string())
.collect();
format!("{} IN ({})", self.dialect.quote(f), vals_str.join(", "))
}
WhereCondition::NotIn(f, vals) => {
let vals_str: Vec<String> = vals
.iter()
.map(|v| v.to_param_with_dialect(&*self.dialect).to_string())
.collect();
format!("{} NOT IN ({})", self.dialect.quote(f), vals_str.join(", "))
}
WhereCondition::Between(f, start, end) => {
format!(
"{} BETWEEN {} AND {}",
self.dialect.quote(f),
start.to_param_with_dialect(&*self.dialect),
end.to_param_with_dialect(&*self.dialect)
)
}
WhereCondition::NotBetween(f, start, end) => {
format!(
"{} NOT BETWEEN {} AND {}",
self.dialect.quote(f),
start.to_param_with_dialect(&*self.dialect),
end.to_param_with_dialect(&*self.dialect)
)
}
WhereCondition::Null(f) => format!("{} IS NULL", self.dialect.quote(f)),
WhereCondition::NotNull(f) => format!("{} IS NOT NULL", self.dialect.quote(f)),
WhereCondition::Exists(s) => format!("EXISTS ({})", s),
WhereCondition::NotExists(s) => format!("NOT EXISTS ({})", s),
})
.collect();
if let Some(sd_cond) = soft_delete_cond {
conditions.push(sd_cond);
}
if let Some(t_cond) = tenant_cond {
conditions.push(t_cond);
}
if let Some(ref cursor) = self.keyset_cursor {
let op = match cursor.direction {
KeysetDirection::After => ">",
KeysetDirection::Before => "<",
};
conditions.push(format!(
"{} {} {}",
self.dialect.quote(&cursor.field),
op,
cursor.value.to_param_with_dialect(&*self.dialect)
));
}
if conditions.is_empty() {
return String::new();
}
let mut groups: Vec<Vec<String>> = Vec::new();
let mut current_group: Vec<String> = Vec::new();
for cond in conditions.iter() {
if let Some(stripped) = cond.strip_prefix("OR ") {
current_group.push(stripped.to_string());
} else {
if !current_group.is_empty() {
groups.push(std::mem::take(&mut current_group));
}
current_group.push(cond.clone());
}
}
if !current_group.is_empty() {
groups.push(current_group);
}
let group_strs: Vec<String> = groups
.iter()
.map(|g| {
if g.len() == 1 {
g[0].clone()
} else {
format!("({})", g.join(" OR "))
}
})
.collect();
format!(" WHERE {}", group_strs.join(" AND "))
}
#[tracing::instrument(skip(self, data), fields(op = "insert"))]
pub fn build_insert(&self, data: &std::collections::HashMap<String, Value>) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if data.is_empty() {
return String::new();
}
let columns: Vec<String> = data.keys().map(|k| self.dialect.quote(k)).collect();
let values: Vec<String> = data
.values()
.map(|v| v.to_param_with_dialect(&*self.dialect).to_string())
.collect();
format!(
"INSERT INTO {} ({}) VALUES ({})",
self.dialect.quote(&table),
columns.join(", "),
values.join(", ")
)
}
#[tracing::instrument(skip(self, data), fields(op = "update"))]
pub fn build_update(&self, data: &std::collections::HashMap<String, Value>) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if data.is_empty() {
return String::new();
}
let set_clauses: Vec<String> = data
.iter()
.map(|(k, v)| {
format!(
"{} = {}",
self.dialect.quote(k),
v.to_param_with_dialect(&*self.dialect)
)
})
.collect();
let mut sql = format!(
"UPDATE {} SET {}",
self.dialect.quote(&table),
set_clauses.join(", ")
);
sql.push_str(&self.build_where_clause());
sql
}
#[tracing::instrument(skip(self), fields(op = "delete"))]
pub fn build_delete(&self) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if let Some(field) = self.soft_delete_field() {
let where_clause = self.build_where_clause();
return format!(
"UPDATE {} SET {} = NOW(){}",
self.dialect.quote(&table),
self.dialect.quote(field),
where_clause
);
}
let mut sql = format!("DELETE FROM {}", self.dialect.quote(&table));
sql.push_str(&self.build_where_clause());
sql
}
pub fn build_force_delete(&self) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!("DELETE FROM {}", self.dialect.quote(&table));
sql.push_str(&self.build_where_clause_with_options(false));
sql
}
fn build_where_clause_with_params(&self) -> (String, Vec<Value>) {
self.build_where_clause_with_params_options(true)
}
fn build_where_clause_with_params_options(
&self,
include_soft_delete: bool,
) -> (String, Vec<Value>) {
let soft_delete_cond = if include_soft_delete {
self.build_soft_delete_condition()
} else {
None
};
let tenant_cond = self.build_tenant_condition();
if self.where_conditions.is_empty()
&& soft_delete_cond.is_none()
&& tenant_cond.is_none()
&& self.keyset_cursor.is_none()
{
return (String::new(), Vec::new());
}
let mut params = Vec::new();
let mut conditions: Vec<String> = self
.where_conditions
.iter()
.map(|cond| match cond {
WhereCondition::And(c) => c.clone(),
WhereCondition::Or(c) => format!("OR {}", c),
WhereCondition::Eq(f, v) => {
params.push(v.clone());
format!("{} = ?", self.dialect.quote(f))
}
WhereCondition::Ne(f, v) => {
params.push(v.clone());
format!("{} != ?", self.dialect.quote(f))
}
WhereCondition::Gt(f, v) => {
params.push(v.clone());
format!("{} > ?", self.dialect.quote(f))
}
WhereCondition::Ge(f, v) => {
params.push(v.clone());
format!("{} >= ?", self.dialect.quote(f))
}
WhereCondition::Lt(f, v) => {
params.push(v.clone());
format!("{} < ?", self.dialect.quote(f))
}
WhereCondition::Le(f, v) => {
params.push(v.clone());
format!("{} <= ?", self.dialect.quote(f))
}
WhereCondition::Like(f, v) => {
params.push(v.clone());
format!("{} LIKE ?", self.dialect.quote(f))
}
WhereCondition::OrEq(f, v) => {
params.push(v.clone());
format!("OR {} = ?", self.dialect.quote(f))
}
WhereCondition::OrNe(f, v) => {
params.push(v.clone());
format!("OR {} != ?", self.dialect.quote(f))
}
WhereCondition::OrGt(f, v) => {
params.push(v.clone());
format!("OR {} > ?", self.dialect.quote(f))
}
WhereCondition::OrGe(f, v) => {
params.push(v.clone());
format!("OR {} >= ?", self.dialect.quote(f))
}
WhereCondition::OrLt(f, v) => {
params.push(v.clone());
format!("OR {} < ?", self.dialect.quote(f))
}
WhereCondition::OrLe(f, v) => {
params.push(v.clone());
format!("OR {} <= ?", self.dialect.quote(f))
}
WhereCondition::OrLike(f, v) => {
params.push(v.clone());
format!("OR {} LIKE ?", self.dialect.quote(f))
}
WhereCondition::In(f, vals) => {
let placeholders: Vec<&str> = vals.iter().map(|_| "?").collect();
params.extend(vals.iter().cloned());
format!("{} IN ({})", self.dialect.quote(f), placeholders.join(", "))
}
WhereCondition::NotIn(f, vals) => {
let placeholders: Vec<&str> = vals.iter().map(|_| "?").collect();
params.extend(vals.iter().cloned());
format!(
"{} NOT IN ({})",
self.dialect.quote(f),
placeholders.join(", ")
)
}
WhereCondition::Between(f, start, end) => {
params.push(start.clone());
params.push(end.clone());
format!("{} BETWEEN ? AND ?", self.dialect.quote(f))
}
WhereCondition::NotBetween(f, start, end) => {
params.push(start.clone());
params.push(end.clone());
format!("{} NOT BETWEEN ? AND ?", self.dialect.quote(f))
}
WhereCondition::Null(f) => format!("{} IS NULL", self.dialect.quote(f)),
WhereCondition::NotNull(f) => format!("{} IS NOT NULL", self.dialect.quote(f)),
WhereCondition::Exists(s) => format!("EXISTS ({})", s),
WhereCondition::NotExists(s) => format!("NOT EXISTS ({})", s),
})
.collect();
if let Some(sd_cond) = soft_delete_cond {
conditions.push(sd_cond);
}
if let Some((t_sql, t_value)) = tenant_cond {
conditions.push(t_sql);
params.push(t_value);
}
if let Some(ref cursor) = self.keyset_cursor {
let op = match cursor.direction {
KeysetDirection::After => ">",
KeysetDirection::Before => "<",
};
conditions.push(format!("{} {} ?", self.dialect.quote(&cursor.field), op));
params.push(cursor.value.clone());
}
if conditions.is_empty() {
return (String::new(), params);
}
let mut groups: Vec<Vec<String>> = Vec::new();
let mut current_group: Vec<String> = Vec::new();
for cond in conditions.iter() {
if let Some(stripped) = cond.strip_prefix("OR ") {
current_group.push(stripped.to_string());
} else {
if !current_group.is_empty() {
groups.push(std::mem::take(&mut current_group));
}
current_group.push(cond.clone());
}
}
if !current_group.is_empty() {
groups.push(current_group);
}
let group_strs: Vec<String> = groups
.iter()
.map(|g| {
if g.len() == 1 {
g[0].clone()
} else {
format!("({})", g.join(" OR "))
}
})
.collect();
(format!(" WHERE {}", group_strs.join(" AND ")), params)
}
pub fn build_select_with_params(&self) -> (String, Vec<Value>) {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let columns = if self.select_columns.is_empty() {
"*".to_string()
} else {
self.select_columns.join(", ")
};
let mut sql = format!("SELECT {} FROM {}", columns, self.dialect.quote(&table));
for join in &self.joins {
match join {
JoinClause::Inner(t, l, r) => {
sql.push_str(&format!(
" INNER JOIN {} ON {} = {}",
self.dialect.quote(t),
self.dialect.quote(l),
self.dialect.quote(r)
));
}
JoinClause::Left(t, l, r) => {
sql.push_str(&format!(
" LEFT JOIN {} ON {} = {}",
self.dialect.quote(t),
self.dialect.quote(l),
self.dialect.quote(r)
));
}
JoinClause::Right(t, l, r) => {
sql.push_str(&format!(
" RIGHT JOIN {} ON {} = {}",
self.dialect.quote(t),
self.dialect.quote(l),
self.dialect.quote(r)
));
}
JoinClause::Cross(t, on) => {
sql.push_str(&format!(
" CROSS JOIN {} ON {}",
self.dialect.quote(t),
self.dialect.quote(on)
));
}
}
}
let mut params = Vec::new();
let (where_clause, where_params) = self.build_where_clause_with_params();
if !where_clause.is_empty() {
sql.push_str(&where_clause);
params = where_params;
}
if !self.group_by.is_empty() {
let cols: Vec<String> = self
.group_by
.iter()
.map(|c| self.dialect.quote(c))
.collect();
sql.push_str(" GROUP BY ");
sql.push_str(&cols.join(", "));
}
if !self.having_conditions.is_empty() {
sql.push_str(" HAVING ");
for (i, cond) in self.having_conditions.iter().enumerate() {
if i > 0 {
sql.push_str(" AND ");
}
if let WhereCondition::And(c) = cond {
sql.push_str(c);
}
}
}
if !self.order_by.is_empty() {
let order_cols: Vec<String> = self
.order_by
.iter()
.map(|o| {
let dir = match o.direction {
OrderDirection::Asc => " ASC",
OrderDirection::Desc => " DESC",
};
format!("{}{}", self.dialect.quote(&o.field), dir)
})
.collect();
sql.push_str(" ORDER BY ");
sql.push_str(&order_cols.join(", "));
}
if let Some(limit) = self.limit_value {
sql.push_str(&format!(" LIMIT {}", limit));
}
if let Some(offset) = self.offset_value {
sql.push_str(&format!(" OFFSET {}", offset));
}
(sql, params)
}
pub fn build_insert_with_params(
&self,
data: &std::collections::HashMap<String, Value>,
) -> (String, Vec<Value>) {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if data.is_empty() {
return (String::new(), Vec::new());
}
let mut columns = Vec::with_capacity(data.len());
let mut params = Vec::with_capacity(data.len());
let placeholders: Vec<&str> = data.iter().map(|_| "?").collect();
for (k, v) in data.iter() {
columns.push(self.dialect.quote(k));
params.push(v.clone());
}
let sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
self.dialect.quote(&table),
columns.join(", "),
placeholders.join(", ")
);
(sql, params)
}
pub fn build_batch_insert_with_params(
&self,
rows: &[std::collections::HashMap<String, Value>],
) -> (String, Vec<Value>) {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if rows.is_empty() {
return (String::new(), Vec::new());
}
let first_row = &rows[0];
let columns: Vec<String> = first_row.keys().cloned().collect();
let quoted_columns: Vec<String> = columns.iter().map(|c| self.dialect.quote(c)).collect();
let mut params = Vec::with_capacity(rows.len() * columns.len());
let mut value_groups: Vec<String> = Vec::with_capacity(rows.len());
for row in rows {
let placeholders: Vec<String> = columns
.iter()
.map(|col| match row.get(col) {
Some(v) => {
params.push(v.clone());
"?".to_string()
}
None => "NULL".to_string(),
})
.collect();
value_groups.push(format!("({})", placeholders.join(", ")));
}
let sql = format!(
"INSERT INTO {} ({}) VALUES {}",
self.dialect.quote(&table),
quoted_columns.join(", "),
value_groups.join(", ")
);
(sql, params)
}
pub fn build_batch_upsert_with_params(
&self,
rows: &[std::collections::HashMap<String, Value>],
conflict_columns: &[&str],
update_columns: &[&str],
) -> Result<(String, Vec<Value>), crate::DbError> {
if rows.is_empty() {
return Err(crate::DbError::InvalidInput(
"build_batch_upsert_with_params: rows cannot be empty".to_string(),
));
}
let (insert_sql, params) = self.build_batch_insert_with_params(rows);
if insert_sql.is_empty() {
return Err(crate::DbError::InvalidInput(
"build_batch_upsert_with_params: failed to build INSERT part".to_string(),
));
}
let all_columns: Vec<String> = rows[0].keys().cloned().collect();
let conflict_clause = self
.dialect
.build_upsert_on_conflict(conflict_columns, update_columns, &all_columns)
.ok_or_else(|| {
crate::DbError::InvalidInput(format!(
"build_batch_upsert_with_params: dialect {:?} does not support upsert (ON CONFLICT / ON DUPLICATE KEY UPDATE). Consider using MERGE statement or individual upserts instead.",
self.dialect.db_type()
))
})?;
let sql = format!("{} {}", insert_sql, conflict_clause);
Ok((sql, params))
}
pub fn build_update_with_params(
&self,
data: &std::collections::HashMap<String, Value>,
) -> (String, Vec<Value>) {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if data.is_empty() {
return (String::new(), Vec::new());
}
let mut set_clauses = Vec::with_capacity(data.len());
let mut params = Vec::with_capacity(data.len());
for (k, v) in data.iter() {
set_clauses.push(format!("{} = ?", self.dialect.quote(k)));
params.push(v.clone());
}
let mut sql = format!(
"UPDATE {} SET {}",
self.dialect.quote(&table),
set_clauses.join(", ")
);
let (where_clause, where_params) = self.build_where_clause_with_params();
if !where_clause.is_empty() {
sql.push_str(&where_clause);
params.extend(where_params);
}
(sql, params)
}
pub fn build_delete_with_params(&self) -> (String, Vec<Value>) {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if let Some(field) = self.soft_delete_field() {
let (where_clause, where_params) = self.build_where_clause_with_params();
let sql = format!(
"UPDATE {} SET {} = NOW(){}",
self.dialect.quote(&table),
self.dialect.quote(field),
where_clause
);
return (sql, where_params);
}
let mut sql = format!("DELETE FROM {}", self.dialect.quote(&table));
let mut params = Vec::new();
let (where_clause, where_params) = self.build_where_clause_with_params();
if !where_clause.is_empty() {
sql.push_str(&where_clause);
params = where_params;
}
(sql, params)
}
pub fn build_force_delete_with_params(&self) -> (String, Vec<Value>) {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!("DELETE FROM {}", self.dialect.quote(&table));
let mut params = Vec::new();
let (where_clause, where_params) = self.build_where_clause_with_params_options(false);
if !where_clause.is_empty() {
sql.push_str(&where_clause);
params = where_params;
}
(sql, params)
}
pub fn build_count(&self) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!(
"SELECT COUNT(*) as total FROM {}",
self.dialect.quote(&table)
);
sql.push_str(&self.build_where_clause());
sql
}
pub fn build_exists(&self) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!("SELECT 1 FROM {}", self.dialect.quote(&table));
sql.push_str(&self.build_where_clause());
sql.push_str(" LIMIT 1");
format!("SELECT EXISTS({})", sql)
}
pub fn build_max(&self, field: &str) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!(
"SELECT MAX({}) as max_val FROM {}",
self.dialect.quote(field),
self.dialect.quote(&table)
);
sql.push_str(&self.build_where_clause());
sql
}
pub fn build_min(&self, field: &str) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!(
"SELECT MIN({}) as min_val FROM {}",
self.dialect.quote(field),
self.dialect.quote(&table)
);
sql.push_str(&self.build_where_clause());
sql
}
pub fn build_sum(&self, field: &str) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!(
"SELECT SUM({}) as sum_val FROM {}",
self.dialect.quote(field),
self.dialect.quote(&table)
);
sql.push_str(&self.build_where_clause());
sql
}
pub fn build_avg(&self, field: &str) -> String {
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
let mut sql = format!(
"SELECT AVG({}) as avg_val FROM {}",
self.dialect.quote(field),
self.dialect.quote(&table)
);
sql.push_str(&self.build_where_clause());
sql
}
pub fn validate(&self) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
let sql = self.build_select();
let mut errors = Vec::new();
if let Err(e) = sz_orm_sql_validator::validate_select(&sql) {
errors.push(e);
}
if !self.joins.is_empty() {
for join in &self.joins {
match join {
JoinClause::Inner(_, left, right)
| JoinClause::Left(_, left, right)
| JoinClause::Right(_, left, right) => {
if let Err(e) = sz_orm_sql_validator::validate_column_name(left) {
errors.push(e);
}
if let Err(e) = sz_orm_sql_validator::validate_column_name(right) {
errors.push(e);
}
}
_ => {}
}
}
}
let table = self
.table
.clone()
.unwrap_or_else(|| M::table_name().to_string());
if let Err(e) = sz_orm_sql_validator::validate_table_name(&table) {
errors.push(e);
}
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
pub fn validate_insert(
&self,
data: &std::collections::HashMap<String, Value>,
) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
let sql = self.build_insert(data);
let mut errors = Vec::new();
if sql.is_empty() {
errors.push(sz_orm_sql_validator::SqlValidationError::EmptyInsertData);
return Err(errors);
}
if let Err(e) = sz_orm_sql_validator::validate_insert(&sql) {
errors.push(e);
}
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
pub fn validate_update(
&self,
data: &std::collections::HashMap<String, Value>,
) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
let sql = self.build_update(data);
let mut errors = Vec::new();
if sql.is_empty() {
errors.push(sz_orm_sql_validator::SqlValidationError::EmptyUpdateData);
return Err(errors);
}
if let Err(e) = sz_orm_sql_validator::validate_update(&sql) {
errors.push(e);
}
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
pub fn validate_delete(&self) -> Result<(), Vec<sz_orm_sql_validator::SqlValidationError>> {
let sql = self.build_delete();
let mut errors = Vec::new();
if let Err(e) = sz_orm_sql_validator::validate_delete(&sql) {
errors.push(e);
}
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
}
impl<M: Model> fmt::Debug for QueryBuilder<M> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("QueryBuilder")
.field("table", &self.table)
.field("select_columns", &self.select_columns)
.field("where_conditions", &self.where_conditions.len())
.field("limit", &self.limit_value)
.finish()
}
}
#[cfg(test)]
#[allow(deprecated)] mod tests {
use super::*;
use crate::DbType;
use crate::get_dialect;
struct TestModel;
impl Model for TestModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"test_models"
}
fn pk(&self) -> Self::PrimaryKey {
1
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
}
#[test]
fn test_query_builder_select() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.select(vec!["id", "name"])
.build_select();
assert!(sql.contains("SELECT id, name FROM"));
assert!(sql.contains("`users`"));
Ok(())
}
#[test]
fn test_query_builder_where() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.where_cond("status = 'active'")
.where_cond("age > 18")
.build_select();
assert!(sql.contains("WHERE"));
assert!(sql.contains("status = 'active'"));
assert!(sql.contains("age > 18"));
Ok(())
}
#[test]
fn test_query_builder_order_by() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.order_by("created_at")
.order_desc("id")
.build_select();
assert!(sql.contains("ORDER BY"));
assert!(sql.contains("`created_at` ASC"));
assert!(sql.contains("`id` DESC"));
Ok(())
}
#[test]
fn test_query_builder_limit_offset() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").limit(10).offset(20).build_select();
assert!(sql.contains("LIMIT 10"));
assert!(sql.contains("OFFSET 20"));
Ok(())
}
#[test]
fn test_query_builder_page() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").page(3, 20).build_select();
assert!(sql.contains("LIMIT 20"));
assert!(sql.contains("OFFSET 40"));
Ok(())
}
#[test]
fn test_query_builder_insert() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let mut data = std::collections::HashMap::new();
data.insert("name".to_string(), Value::String("test".to_string()));
data.insert("age".to_string(), Value::I64(25));
let sql = builder.table("users").build_insert(&data);
assert!(sql.contains("INSERT INTO"));
assert!(sql.contains("`name`"));
assert!(sql.contains("'test'"));
Ok(())
}
#[test]
fn test_query_builder_update() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let mut data = std::collections::HashMap::new();
data.insert("name".to_string(), Value::String("updated".to_string()));
let sql = builder
.table("users")
.where_cond("id = 1")
.build_update(&data);
assert!(sql.contains("UPDATE"));
assert!(sql.contains("`name` = 'updated'"));
assert!(sql.contains("WHERE"));
Ok(())
}
#[test]
fn test_query_builder_delete() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").where_cond("id = 1").build_delete();
assert!(sql.contains("DELETE FROM"));
assert!(sql.contains("WHERE"));
Ok(())
}
#[test]
fn test_query_builder_count() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").build_count();
assert!(sql.contains("SELECT COUNT(*)"));
assert!(sql.contains("FROM"));
Ok(())
}
#[test]
fn test_query_builder_where_in() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.where_in("id", vec![Value::I64(1), Value::I64(2), Value::I64(3)])
.build_select();
assert!(sql.contains("IN ("));
Ok(())
}
#[test]
fn test_query_builder_where_between() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.where_between("age", Value::I64(18), Value::I64(30))
.build_select();
assert!(sql.contains("BETWEEN"));
Ok(())
}
#[test]
fn test_query_builder_where_null() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.where_null("deleted_at")
.build_select();
assert!(sql.contains("IS NULL"));
Ok(())
}
#[test]
fn test_query_builder_join() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder
.table("users")
.join_inner("posts", "users.id", "posts.user_id")
.build_select();
assert!(sql.contains("INNER JOIN"));
assert!(sql.contains("`posts`"));
Ok(())
}
#[test]
fn test_query_builder_group_by() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").group_by("status").build_select();
assert!(sql.contains("GROUP BY"));
assert!(sql.contains("`status`"));
Ok(())
}
#[test]
fn test_query_builder_max() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").build_max("score");
assert!(sql.contains("MAX("));
assert!(sql.contains("`score`"));
Ok(())
}
#[test]
fn test_query_builder_min() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("users").build_min("price");
assert!(sql.contains("MIN("));
assert!(sql.contains("`price`"));
Ok(())
}
#[test]
fn test_query_builder_sum() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("orders").build_sum("amount");
assert!(sql.contains("SUM("));
assert!(sql.contains("`amount`"));
Ok(())
}
#[test]
fn test_query_builder_avg() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let sql = builder.table("scores").build_avg("value");
assert!(sql.contains("AVG("));
assert!(sql.contains("`value`"));
Ok(())
}
#[test]
fn test_validator_select() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder.table("users").select(vec!["id", "name"]).validate();
assert!(result.is_ok());
Ok(())
}
#[test]
fn test_validator_select_with_join() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder
.table("users")
.join_inner("posts", "users.id", "posts.user_id")
.validate();
assert!(result.is_ok());
Ok(())
}
#[test]
fn test_validator_insert() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let mut data = std::collections::HashMap::new();
data.insert("name".to_string(), Value::String("test".to_string()));
let result = builder.table("users").validate_insert(&data);
assert!(result.is_ok());
Ok(())
}
#[test]
fn test_validator_insert_empty_data() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let data = std::collections::HashMap::new();
let result = builder.table("users").validate_insert(&data);
assert!(result.is_err());
Ok(())
}
#[test]
fn test_validator_update() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let mut data = std::collections::HashMap::new();
data.insert("name".to_string(), Value::String("updated".to_string()));
let result = builder.table("users").validate_update(&data);
assert!(result.is_ok());
Ok(())
}
#[test]
fn test_validator_update_empty_data() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let data = std::collections::HashMap::new();
let result = builder.table("users").validate_update(&data);
assert!(result.is_err());
Ok(())
}
#[test]
fn test_validator_delete() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder
.table("users")
.where_cond("id = 1")
.validate_delete();
assert!(result.is_ok());
Ok(())
}
#[test]
fn test_validator_delete_no_where() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder.table("users").validate_delete();
assert!(result.is_ok());
Ok(())
}
#[test]
fn test_m3_select_quoted_valid_columns() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let builder = builder.table("users").select_quoted(vec!["id", "name"])?;
let sql = builder.build_select();
assert!(sql.contains("SELECT `id`, `name` FROM"));
assert!(sql.contains("`users`"));
Ok(())
}
#[test]
fn test_m3_select_quoted_rejects_sql_injection() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder
.table("users")
.select_quoted(vec!["id; DROP TABLE users"]);
assert!(result.is_err());
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder.table("users").select_quoted(vec!["name'"]);
assert!(result.is_err());
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder.table("users").select_quoted(vec!["1col"]);
assert!(result.is_err());
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let result = builder.table("users").select_quoted(vec!["col name"]);
assert!(result.is_err());
Ok(())
}
#[test]
fn test_m3_select_quoted_postgresql_dialect() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::PostgreSQL)?;
let builder = QueryBuilder::<TestModel>::new(dialect);
let builder = builder.table("users").select_quoted(vec!["id", "name"])?;
let sql = builder.build_select();
assert!(sql.contains("SELECT \"id\", \"name\" FROM"));
assert!(sql.contains("\"users\""));
Ok(())
}
struct SoftDeleteModel;
impl Model for SoftDeleteModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"soft_users"
}
fn pk(&self) -> Self::PrimaryKey {
1
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
fn soft_delete_field() -> Option<&'static str> {
Some("deleted_at")
}
}
#[test]
fn test_p01_soft_delete_select_auto_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<SoftDeleteModel>::new(dialect);
let sql = builder.table("soft_users").build_select();
assert!(
sql.contains("`deleted_at` IS NULL"),
"软删除模型 SELECT 必须自动追加 `deleted_at` IS NULL,实际: {}",
sql
);
Ok(())
}
#[test]
fn test_p01_soft_delete_select_with_user_where() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.where_eq("status", Value::String("active".into()))
.build_select();
assert!(sql.contains("`status` = "), "用户条件应保留: {}", sql);
assert!(
sql.contains("`deleted_at` IS NULL"),
"软删除条件应自动追加: {}",
sql
);
Ok(())
}
#[test]
fn test_p01_soft_delete_without_soft_delete() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.without_soft_delete()
.build_select();
assert!(
!sql.contains("`deleted_at` IS NULL"),
"without_soft_delete 应禁用过滤,实际: {}",
sql
);
assert!(
!sql.contains("WHERE"),
"无用户条件 + 禁用软删除应无 WHERE 子句: {}",
sql
);
Ok(())
}
#[test]
fn test_p01_soft_delete_delete_becomes_update() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.where_eq("id", Value::I64(42))
.build_delete();
assert!(
sql.starts_with("UPDATE"),
"软删除模型的 build_delete 应生成 UPDATE,实际: {}",
sql
);
assert!(
!sql.contains("DELETE FROM"),
"不应生成 DELETE FROM: {}",
sql
);
assert!(
sql.contains("`deleted_at` = NOW()"),
"应设置 deleted_at = NOW(): {}",
sql
);
assert!(
sql.contains("`deleted_at` IS NULL"),
"软删除 UPDATE 应追加 deleted_at IS NULL 防止重复删除: {}",
sql
);
Ok(())
}
#[test]
fn test_p01_soft_delete_force_delete() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.where_eq("id", Value::I64(99))
.build_force_delete();
assert!(
sql.starts_with("DELETE FROM"),
"build_force_delete 应生成 DELETE FROM,实际: {}",
sql
);
assert!(
!sql.contains("`deleted_at` IS NULL"),
"物理删除不应追加软删除过滤: {}",
sql
);
Ok(())
}
#[test]
fn test_p01_soft_delete_select_with_params() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.where_eq("id", Value::I64(1))
.build_select_with_params();
assert!(
sql.contains("`deleted_at` IS NULL"),
"参数化版本也应自动追加软删除: {}",
sql
);
assert_eq!(params.len(), 1, "参数应为 1 个(用户 where_eq 的值)");
assert_eq!(params[0], Value::I64(1));
Ok(())
}
#[test]
fn test_p01_soft_delete_delete_with_params_becomes_update() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.where_eq("id", Value::I64(7))
.build_delete_with_params();
assert!(sql.starts_with("UPDATE"), "应生成 UPDATE: {}", sql);
assert!(
sql.contains("`deleted_at` = NOW()"),
"应设置 NOW(): {}",
sql
);
assert_eq!(params.len(), 1, "参数应为 1 个(WHERE 的值)");
Ok(())
}
#[test]
fn test_p01_soft_delete_force_delete_with_params() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.where_eq("id", Value::I64(11))
.build_force_delete_with_params();
assert!(sql.starts_with("DELETE FROM"), "应生成 DELETE: {}", sql);
assert!(
!sql.contains("`deleted_at` IS NULL"),
"不应追加软删除过滤: {}",
sql
);
assert_eq!(params.len(), 1);
Ok(())
}
#[test]
fn test_p01_non_soft_delete_model_unchanged() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("id", Value::I64(1))
.build_select();
assert!(
!sql.contains("deleted_at"),
"非软删除模型不应追加 deleted_at: {}",
sql
);
let dialect = get_dialect(DbType::MySQL)?;
let del_sql = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("id", Value::I64(1))
.build_delete();
assert!(
del_sql.starts_with("DELETE FROM"),
"非软删除模型 build_delete 应生成 DELETE: {}",
del_sql
);
Ok(())
}
#[test]
fn test_p01_soft_delete_count_auto_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<SoftDeleteModel>::new(dialect)
.table("soft_users")
.build_count();
assert!(
sql.contains("`deleted_at` IS NULL"),
"build_count 也应追加软删除过滤: {}",
sql
);
Ok(())
}
#[test]
fn test_p02_where_eq_uses_placeholder() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("name", Value::String("alice".into()))
.build_select_with_params();
assert!(sql.contains("`name` = ?"), "应使用 ? 占位符: {}", sql);
assert!(!sql.contains("'alice'"), "不应内嵌值到 SQL: {}", sql);
assert_eq!(params.len(), 1);
assert_eq!(params[0], Value::String("alice".into()));
Ok(())
}
#[test]
fn test_p02_where_like_uses_placeholder() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_like("name", Value::String("%alice%".into()))
.build_select_with_params();
assert!(sql.contains("`name` LIKE ?"), "应使用 LIKE ?: {}", sql);
assert!(!sql.contains("%alice%"), "不应内嵌 pattern: {}", sql);
assert_eq!(params.len(), 1);
Ok(())
}
#[test]
fn test_p02_where_ne_uses_placeholder() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_ne("status", Value::I64(0))
.build_select_with_params();
assert!(sql.contains("`status` != ?"), "应使用 != ?: {}", sql);
assert!(!sql.contains("!= 0"), "不应内嵌值: {}", sql);
assert_eq!(params.len(), 1);
assert_eq!(params[0], Value::I64(0));
Ok(())
}
#[test]
fn test_p02_where_ge_uses_placeholder() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_ge("age", Value::I64(18))
.build_select_with_params();
assert!(sql.contains("`age` >= ?"), "应使用 >= ?: {}", sql);
assert!(!sql.contains(">= 18"), "不应内嵌值: {}", sql);
assert_eq!(params.len(), 1);
assert_eq!(params[0], Value::I64(18));
Ok(())
}
#[test]
fn test_p02_where_lt_uses_placeholder() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_lt("score", Value::F64(60.0))
.build_select_with_params();
assert!(sql.contains("`score` < ?"), "应使用 < ?: {}", sql);
assert!(!sql.contains("< 60"), "不应内嵌值: {}", sql);
assert_eq!(params.len(), 1);
assert_eq!(params[0], Value::F64(60.0));
Ok(())
}
#[test]
fn test_p02_injection_protection_drop_table() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let evil_input = "'; DROP TABLE users; --".to_string();
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("name", Value::String(evil_input.clone()))
.build_select_with_params();
assert!(!sql.contains("DROP TABLE"), "SQL 注入未防护: {}", sql);
assert_eq!(params.len(), 1);
assert_eq!(params[0], Value::String(evil_input));
assert_eq!(sql.matches('?').count(), 1);
Ok(())
}
#[test]
fn test_p02_injection_protection_or_one_equals_one() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let evil = "' OR '1'='1".to_string();
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("name", Value::String(evil.clone()))
.build_select_with_params();
assert!(!sql.contains("OR '1'='1'"), "OR 1=1 注入未防护: {}", sql);
assert_eq!(params.len(), 1);
assert_eq!(params[0], Value::String(evil));
Ok(())
}
#[test]
fn test_p02_multiple_params_order() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("name", Value::String("alice".into()))
.where_gt("age", Value::I64(18))
.where_le("score", Value::F64(99.5))
.build_select_with_params();
assert_eq!(sql.matches('?').count(), 3, "应有 3 个占位符: {}", sql);
assert_eq!(params.len(), 3);
assert_eq!(params[0], Value::String("alice".into()));
assert_eq!(params[1], Value::I64(18));
assert_eq!(params[2], Value::F64(99.5));
Ok(())
}
#[test]
fn test_p02_where_in_uses_placeholders() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_in("id", vec![Value::I64(1), Value::I64(2), Value::I64(3)])
.build_select_with_params();
assert!(
sql.contains("`id` IN (?, ?, ?)"),
"应使用 3 个占位符: {}",
sql
);
assert_eq!(params.len(), 3);
Ok(())
}
#[test]
fn test_p02_where_between_uses_placeholders() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_between("age", Value::I64(18), Value::I64(65))
.build_select_with_params();
assert!(
sql.contains("`age` BETWEEN ? AND ?"),
"应使用 2 个占位符: {}",
sql
);
assert_eq!(params.len(), 2);
assert_eq!(params[0], Value::I64(18));
assert_eq!(params[1], Value::I64(65));
Ok(())
}
#[test]
fn test_p02_update_params_order_set_before_where() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let mut data = std::collections::HashMap::new();
data.insert("name".to_string(), Value::String("bob".into()));
data.insert("age".to_string(), Value::I64(30));
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("id", Value::I64(99))
.build_update_with_params(&data);
assert_eq!(sql.matches('?').count(), 3, "应有 3 个 ?: {}", sql);
assert_eq!(params.len(), 3);
assert_eq!(params[2], Value::I64(99));
Ok(())
}
#[test]
fn test_p02_build_where_clause_inlines_value() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.where_eq("name", Value::String("alice".into()))
.build_select();
assert!(
sql.contains("`name` = "),
"无参数版本应含 WHERE 条件: {}",
sql
);
assert!(
!sql.contains("`name` = ?"),
"无参数版本不应使用 ? 占位符: {}",
sql
);
Ok(())
}
#[test]
fn test_p01_is_soft_delete_disabled_flag() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<SoftDeleteModel>::new(dialect);
assert!(!builder.is_soft_delete_disabled(), "默认应启用软删除过滤");
let builder =
QueryBuilder::<SoftDeleteModel>::new(get_dialect(DbType::MySQL)?).without_soft_delete();
assert!(
builder.is_soft_delete_disabled(),
"without_soft_delete 后应反映禁用状态"
);
Ok(())
}
struct TenantModel;
impl Model for TenantModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"orders"
}
fn pk(&self) -> Self::PrimaryKey {
1
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
fn tenant_field() -> Option<&'static str> {
Some("tenant_id")
}
}
struct SoftDeleteAndTenantModel;
impl Model for SoftDeleteAndTenantModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"documents"
}
fn pk(&self) -> Self::PrimaryKey {
1
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
fn soft_delete_field() -> Option<&'static str> {
Some("deleted_at")
}
fn tenant_field() -> Option<&'static str> {
Some("tenant_id")
}
}
#[test]
fn test_p03_tenant_select_auto_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(42)
.build_select_with_params();
assert!(
sql.contains("`tenant_id` = ?"),
"多租户模型应自动追加 tenant_id = ?: {}",
sql
);
assert_eq!(params.len(), 1, "应有 1 个参数(tenant_id 值)");
assert_eq!(params[0], Value::I64(42));
Ok(())
}
#[test]
fn test_p03_tenant_select_with_user_where() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(7)
.where_eq("status", Value::String("active".into()))
.build_select_with_params();
assert!(sql.contains("`status` = ?"), "用户条件应保留: {}", sql);
assert!(
sql.contains("`tenant_id` = ?"),
"租户条件应自动追加: {}",
sql
);
assert_eq!(params.len(), 2, "应有 2 个参数");
assert_eq!(params[0], Value::String("active".into()));
assert_eq!(params[1], Value::I64(7));
Ok(())
}
#[test]
fn test_p03_tenant_without_tenant() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(42)
.without_tenant()
.build_select_with_params();
assert!(
!sql.contains("`tenant_id` = ?"),
"without_tenant 应禁用过滤: {}",
sql
);
assert_eq!(params.len(), 0, "不应有租户参数");
Ok(())
}
#[test]
fn test_p03_tenant_delete_auto_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(99)
.where_eq("id", Value::I64(1))
.build_delete_with_params();
assert!(
sql.contains("`tenant_id` = ?"),
"删除应自动追加租户条件: {}",
sql
);
assert_eq!(params.len(), 2);
assert_eq!(params[0], Value::I64(1));
assert_eq!(params[1], Value::I64(99));
Ok(())
}
#[test]
fn test_p03_tenant_update_auto_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let mut data = std::collections::HashMap::new();
data.insert("status".to_string(), Value::String("shipped".into()));
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(5)
.where_eq("id", Value::I64(10))
.build_update_with_params(&data);
assert!(
sql.contains("`tenant_id` = ?"),
"更新应自动追加租户条件: {}",
sql
);
assert_eq!(params.len(), 3);
assert_eq!(params[2], Value::I64(5));
Ok(())
}
#[test]
fn test_p03_tenant_count_auto_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let sql = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(42)
.build_count();
assert!(
sql.contains("`tenant_id` = 42"),
"build_count 应追加租户条件(无参数版本内嵌值): {}",
sql
);
Ok(())
}
#[test]
fn test_p03_non_tenant_model_unchanged() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TestModel>::new(dialect)
.table("users")
.with_tenant_id(42)
.build_select_with_params();
assert!(
!sql.contains("tenant_id"),
"非多租户模型不应追加 tenant_id: {}",
sql
);
assert_eq!(params.len(), 0);
Ok(())
}
#[test]
fn test_p03_tenant_no_id_no_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.build_select_with_params();
assert!(
!sql.contains("tenant_id"),
"未设置 tenant_id 时不应追加过滤: {}",
sql
);
assert_eq!(params.len(), 0);
Ok(())
}
#[test]
fn test_p03_soft_delete_and_tenant_combined() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<SoftDeleteAndTenantModel>::new(dialect)
.table("documents")
.with_tenant_id(100)
.where_eq("title", Value::String("report".into()))
.build_select_with_params();
assert!(
sql.contains("`deleted_at` IS NULL"),
"应追加软删除条件: {}",
sql
);
assert!(sql.contains("`tenant_id` = ?"), "应追加租户条件: {}", sql);
assert!(sql.contains("`title` = ?"), "用户条件应保留: {}", sql);
assert_eq!(params.len(), 2);
assert_eq!(params[0], Value::String("report".into()));
assert_eq!(params[1], Value::I64(100));
Ok(())
}
#[test]
fn test_p03_without_tenant_and_soft_delete() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<SoftDeleteAndTenantModel>::new(dialect)
.table("documents")
.with_tenant_id(100)
.without_tenant()
.without_soft_delete()
.build_select_with_params();
assert!(
!sql.contains("`deleted_at` IS NULL"),
"应禁用软删除: {}",
sql
);
assert!(!sql.contains("`tenant_id` = ?"), "应禁用租户: {}", sql);
assert_eq!(params.len(), 0);
Ok(())
}
#[test]
fn test_p03_is_tenant_disabled_flag() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let builder = QueryBuilder::<TenantModel>::new(dialect);
assert!(!builder.is_tenant_disabled(), "默认应启用租户过滤");
let builder = QueryBuilder::<TenantModel>::new(get_dialect(DbType::MySQL)?)
.with_tenant_id(1)
.without_tenant();
assert!(
builder.is_tenant_disabled(),
"without_tenant 后应反映禁用状态"
);
Ok(())
}
#[test]
fn test_p03_tenant_force_delete_keeps_tenant_filter() -> Result<(), crate::DbError> {
let dialect = get_dialect(DbType::MySQL)?;
let (sql, params) = QueryBuilder::<TenantModel>::new(dialect)
.table("orders")
.with_tenant_id(42)
.where_eq("id", Value::I64(999))
.build_force_delete_with_params();
assert!(
sql.contains("`tenant_id` = ?"),
"物理删除应保留租户条件: {}",
sql
);
assert_eq!(params.len(), 2);
assert_eq!(params[0], Value::I64(999));
assert_eq!(params[1], Value::I64(42));
Ok(())
}
}