use sz_orm_core::DbType;
fn quote_ident(s: &str) -> String {
s.split('.')
.map(|part| format!("`{}`", part.replace('`', "``")))
.collect::<Vec<_>>()
.join(".")
}
fn check_where_injection(condition: &str) {
let upper = condition.to_uppercase();
const SQL_KEYWORDS: &[&str] = &[
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE", "EXEC", "CREATE", "GRANT",
"REVOKE",
];
for kw in SQL_KEYWORDS {
let pattern1 = format!(";{}", kw);
let pattern2 = format!("; {}", kw);
if upper.contains(&pattern1) || upper.contains(&pattern2) {
panic!(
"SQL injection detected in where_clause: semicolon followed by {} keyword: {:?}",
kw, condition
);
}
}
if condition.contains("--") {
panic!(
"SQL injection detected in where_clause: line comment '--' not allowed: {:?}",
condition
);
}
if condition.contains("/*") || condition.contains("*/") {
panic!(
"SQL injection detected in where_clause: block comment '/*' or '*/' not allowed: {:?}",
condition
);
}
}
pub struct Query;
impl Query {
pub fn select() -> SelectQuery {
SelectQuery::new()
}
pub fn insert() -> InsertQuery {
InsertQuery::new()
}
pub fn update() -> UpdateQuery {
UpdateQuery::new()
}
pub fn delete() -> DeleteQuery {
DeleteQuery::new()
}
}
#[derive(Debug, Clone, Default)]
pub struct SelectQuery {
columns: Vec<String>,
from_table: Option<String>,
joins: Vec<String>,
wheres: Vec<String>,
order_by: Vec<String>,
group_by: Vec<String>,
having: Vec<String>,
limit: Option<u64>,
offset: Option<u64>,
distinct: bool,
}
impl SelectQuery {
pub fn new() -> Self {
Self::default()
}
pub fn distinct(mut self) -> Self {
self.distinct = true;
self
}
pub fn column(mut self, name: &str) -> Self {
self.columns.push(name.to_string());
self
}
pub fn columns(mut self, names: &[&str]) -> Self {
for n in names {
self.columns.push(n.to_string());
}
self
}
pub fn all_columns(self) -> Self {
self.column("*")
}
pub fn from(mut self, table: &str) -> Self {
self.from_table = Some(table.to_string());
self
}
pub fn inner_join(mut self, table: &str, on: &str) -> Self {
self.joins.push(format!(
"INNER JOIN {} ON {}",
Self::quote_join_table(table),
on
));
self
}
pub fn left_join(mut self, table: &str, on: &str) -> Self {
self.joins.push(format!(
"LEFT JOIN {} ON {}",
Self::quote_join_table(table),
on
));
self
}
pub fn right_join(mut self, table: &str, on: &str) -> Self {
self.joins.push(format!(
"RIGHT JOIN {} ON {}",
Self::quote_join_table(table),
on
));
self
}
fn quote_join_table(table: &str) -> String {
if let Some((tbl, alias)) = table.rsplit_once(' ') {
if alias.to_uppercase() == "AS" {
format!("{} AS {}", quote_ident(tbl), alias)
} else {
format!("{} {}", quote_ident(tbl), alias)
}
} else {
quote_ident(table)
}
}
pub fn where_clause(mut self, condition: &str) -> Self {
check_where_injection(condition);
self.wheres.push(condition.to_string());
self
}
pub fn or_where(mut self, condition: &str) -> Self {
check_where_injection(condition);
self.wheres.push(format!("OR {}", condition));
self
}
pub fn group_by(mut self, column: &str) -> Self {
self.group_by.push(column.to_string());
self
}
pub fn having(mut self, condition: &str) -> Self {
self.having.push(condition.to_string());
self
}
pub fn order_by(mut self, column: &str, asc: bool) -> Self {
let dir = if asc { "ASC" } else { "DESC" };
self.order_by.push(format!("{} {}", column, dir));
self
}
pub fn limit(mut self, n: u64) -> Self {
self.limit = Some(n);
self
}
pub fn offset(mut self, n: u64) -> Self {
self.offset = Some(n);
self
}
pub fn paginate(self, page: u64, size: u64) -> Self {
let offset = (page.saturating_sub(1)) * size;
self.limit(size).offset(offset)
}
pub fn build(self, db_type: DbType) -> String {
let dialect = match sz_orm_core::get_dialect(db_type) {
Ok(d) => d,
Err(_) => return String::new(),
};
let mut sql = String::new();
sql.push_str("SELECT ");
if self.distinct {
sql.push_str("DISTINCT ");
}
if self.columns.is_empty() {
sql.push('*');
} else {
let cols: Vec<String> = self
.columns
.iter()
.map(|c| {
if c == "*" {
c.clone()
} else {
dialect.quote(c)
}
})
.collect();
sql.push_str(&cols.join(", "));
}
if let Some(table) = self.from_table {
sql.push_str(" FROM ");
sql.push_str(&dialect.quote(&table));
}
for join in &self.joins {
sql.push(' ');
sql.push_str(join);
}
if !self.wheres.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&self.wheres[0]);
for w in &self.wheres[1..] {
if w.starts_with("OR ") {
sql.push(' ');
sql.push_str(w);
} else {
sql.push_str(" AND ");
sql.push_str(w);
}
}
}
if !self.group_by.is_empty() {
sql.push_str(" GROUP BY ");
sql.push_str(
&self
.group_by
.iter()
.map(|c| quote_ident(c))
.collect::<Vec<_>>()
.join(", "),
);
}
if !self.having.is_empty() {
sql.push_str(" HAVING ");
sql.push_str(&self.having.join(" AND "));
}
if !self.order_by.is_empty() {
sql.push_str(" ORDER BY ");
sql.push_str(
&self
.order_by
.iter()
.map(|s| {
if let Some((col, dir)) = s.rsplit_once(' ') {
format!("{} {}", quote_ident(col), dir)
} else {
quote_ident(s)
}
})
.collect::<Vec<_>>()
.join(", "),
);
}
if let Some(limit) = self.limit {
sql.push_str(&format!(" LIMIT {}", limit));
}
if let Some(offset) = self.offset {
sql.push_str(&format!(" OFFSET {}", offset));
}
sql
}
}
#[derive(Debug, Clone, Default)]
pub struct InsertQuery {
table: Option<String>,
columns: Vec<String>,
values: Vec<String>,
}
impl InsertQuery {
pub fn new() -> Self {
Self::default()
}
pub fn into_table(mut self, table: &str) -> Self {
self.table = Some(table.to_string());
self
}
pub fn value(mut self, column: &str, value: &str) -> Self {
self.columns.push(column.to_string());
self.values.push(value.to_string());
self
}
pub fn values(mut self, pairs: &[(&str, &str)]) -> Self {
for (c, v) in pairs {
self.columns.push(c.to_string());
self.values.push(v.to_string());
}
self
}
pub fn build(self) -> String {
let table = self.table.unwrap_or_default();
if table.is_empty() || self.columns.is_empty() {
return String::new();
}
let cols: Vec<String> = self.columns.iter().map(|c| quote_ident(c)).collect();
let vals: Vec<String> = self.values.iter().map(|v| v.to_string()).collect();
format!(
"INSERT INTO {} ({}) VALUES ({})",
quote_ident(&table),
cols.join(", "),
vals.join(", ")
)
}
pub fn build_with_dialect(self, db_type: DbType) -> String {
let dialect = match sz_orm_core::get_dialect(db_type) {
Ok(d) => d,
Err(_) => return String::new(),
};
let table = self.table.unwrap_or_default();
if table.is_empty() || self.columns.is_empty() {
return String::new();
}
let cols: Vec<String> = self.columns.iter().map(|c| dialect.quote(c)).collect();
format!(
"INSERT INTO {} ({}) VALUES ({})",
dialect.quote(&table),
cols.join(", "),
self.values.join(", ")
)
}
}
#[derive(Debug, Clone, Default)]
pub struct UpdateQuery {
table: Option<String>,
sets: Vec<(String, String)>,
wheres: Vec<String>,
}
impl UpdateQuery {
pub fn new() -> Self {
Self::default()
}
pub fn table(mut self, table: &str) -> Self {
self.table = Some(table.to_string());
self
}
pub fn set(mut self, column: &str, value: &str) -> Self {
self.sets.push((column.to_string(), value.to_string()));
self
}
pub fn sets(mut self, pairs: &[(&str, &str)]) -> Self {
for (c, v) in pairs {
self.sets.push((c.to_string(), v.to_string()));
}
self
}
pub fn where_clause(mut self, condition: &str) -> Self {
check_where_injection(condition);
self.wheres.push(condition.to_string());
self
}
pub fn build(self) -> String {
let table = self.table.unwrap_or_default();
if table.is_empty() || self.sets.is_empty() {
return String::new();
}
let set_str: Vec<String> = self
.sets
.iter()
.map(|(c, v)| format!("{} = {}", quote_ident(c), v))
.collect();
let mut sql = format!("UPDATE {} SET {}", quote_ident(&table), set_str.join(", "));
if !self.wheres.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&self.wheres.join(" AND "));
}
sql
}
pub fn build_with_dialect(self, db_type: DbType) -> String {
let dialect = match sz_orm_core::get_dialect(db_type) {
Ok(d) => d,
Err(_) => return String::new(),
};
let table = self.table.unwrap_or_default();
if table.is_empty() || self.sets.is_empty() {
return String::new();
}
let set_str: Vec<String> = self
.sets
.iter()
.map(|(c, v)| format!("{} = {}", dialect.quote(c), v))
.collect();
let mut sql = format!(
"UPDATE {} SET {}",
dialect.quote(&table),
set_str.join(", ")
);
if !self.wheres.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&self.wheres.join(" AND "));
}
sql
}
}
#[derive(Debug, Clone, Default)]
pub struct DeleteQuery {
table: Option<String>,
wheres: Vec<String>,
}
impl DeleteQuery {
pub fn new() -> Self {
Self::default()
}
pub fn from_table(mut self, table: &str) -> Self {
self.table = Some(table.to_string());
self
}
pub fn where_clause(mut self, condition: &str) -> Self {
check_where_injection(condition);
self.wheres.push(condition.to_string());
self
}
pub fn build(self) -> String {
let table = self.table.unwrap_or_default();
if table.is_empty() {
return String::new();
}
let mut sql = format!("DELETE FROM {}", quote_ident(&table));
if !self.wheres.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&self.wheres.join(" AND "));
}
sql
}
pub fn build_with_dialect(self, db_type: DbType) -> String {
let dialect = match sz_orm_core::get_dialect(db_type) {
Ok(d) => d,
Err(_) => return String::new(),
};
let table = self.table.unwrap_or_default();
if table.is_empty() {
return String::new();
}
let mut sql = format!("DELETE FROM {}", dialect.quote(&table));
if !self.wheres.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&self.wheres.join(" AND "));
}
sql
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_select_basic() {
let sql = Query::select()
.column("id")
.column("name")
.from("users")
.build(DbType::MySQL);
assert!(sql.starts_with("SELECT "));
assert!(sql.contains("`id`"));
assert!(sql.contains("`name`"));
assert!(sql.contains("FROM `users`"));
}
#[test]
fn test_select_star() {
let sql = Query::select()
.all_columns()
.from("users")
.build(DbType::MySQL);
assert!(sql.contains("SELECT *"));
assert!(sql.contains("FROM `users`"));
}
#[test]
fn test_select_distinct() {
let sql = Query::select()
.distinct()
.column("name")
.from("users")
.build(DbType::MySQL);
assert!(sql.contains("SELECT DISTINCT"));
}
#[test]
fn test_select_with_where() {
let sql = Query::select()
.column("id")
.from("users")
.where_clause("age > 18")
.where_clause("status = 'active'")
.build(DbType::MySQL);
assert!(sql.contains("WHERE age > 18 AND status = 'active'"));
}
#[test]
fn test_select_with_or_where() {
let sql = Query::select()
.column("id")
.from("users")
.where_clause("age > 18")
.or_where("role = 'admin'")
.build(DbType::MySQL);
assert!(sql.contains("WHERE age > 18 OR role = 'admin'"));
}
#[test]
fn test_select_with_inner_join() {
let sql = Query::select()
.column("u.id")
.from("users u")
.inner_join("orders o", "u.id = o.user_id")
.build(DbType::MySQL);
assert!(sql.contains("INNER JOIN `orders` o ON u.id = o.user_id"));
}
#[test]
fn test_select_with_left_join() {
let sql = Query::select()
.column("u.id")
.from("users u")
.left_join("profiles p", "u.id = p.user_id")
.build(DbType::MySQL);
assert!(sql.contains("LEFT JOIN `profiles` p ON u.id = p.user_id"));
}
#[test]
fn test_select_with_order_by() {
let sql = Query::select()
.column("id")
.from("users")
.order_by("created_at", true)
.order_by("id", false)
.build(DbType::MySQL);
assert!(sql.contains("ORDER BY `created_at` ASC, `id` DESC"));
}
#[test]
fn test_select_with_limit_offset() {
let sql = Query::select()
.column("id")
.from("users")
.limit(10)
.offset(20)
.build(DbType::MySQL);
assert!(sql.contains("LIMIT 10"));
assert!(sql.contains("OFFSET 20"));
}
#[test]
fn test_select_paginate() {
let sql = Query::select()
.column("id")
.from("users")
.paginate(3, 20)
.build(DbType::MySQL);
assert!(sql.contains("LIMIT 20"));
assert!(sql.contains("OFFSET 40"));
}
#[test]
fn test_select_with_group_by_having() {
let sql = Query::select()
.column("status")
.from("users")
.group_by("status")
.having("COUNT(*) > 5")
.build(DbType::MySQL);
assert!(sql.contains("GROUP BY `status`"));
assert!(sql.contains("HAVING COUNT(*) > 5"));
}
#[test]
fn test_select_postgres_dialect() {
let sql = Query::select()
.column("id")
.from("users")
.build(DbType::PostgreSQL);
assert!(sql.contains("\"id\""));
assert!(sql.contains("FROM \"users\""));
}
#[test]
fn test_select_sqlite_dialect() {
let sql = Query::select()
.column("id")
.from("users")
.build(DbType::Sqlite);
assert!(sql.contains("\"id\""));
}
#[test]
fn test_select_multiple_joins() {
let sql = Query::select()
.column("u.id")
.from("users u")
.inner_join("orders o", "u.id = o.user_id")
.left_join("profiles p", "u.id = p.user_id")
.build(DbType::MySQL);
assert!(sql.contains("INNER JOIN `orders` o"));
assert!(sql.contains("LEFT JOIN `profiles` p"));
}
#[test]
fn test_select_columns_multiple() {
let sql = Query::select()
.columns(&["id", "name", "email"])
.from("users")
.build(DbType::MySQL);
assert!(sql.contains("`id`, `name`, `email`"));
}
#[test]
fn test_select_no_columns_defaults_star() {
let sql = Query::select().from("users").build(DbType::MySQL);
assert!(sql.contains("SELECT *"));
}
#[test]
fn test_insert_basic() {
let sql = Query::insert()
.into_table("users")
.value("name", "'Alice'")
.value("age", "30")
.build();
assert!(sql.starts_with("INSERT INTO `users`"));
assert!(sql.contains("`name`, `age`"));
assert!(sql.contains("'Alice', 30"));
}
#[test]
fn test_insert_values_batch() {
let sql = Query::insert()
.into_table("users")
.values(&[("name", "'Bob'"), ("age", "25"), ("email", "'bob@x.com'")])
.build();
assert!(sql.contains("`name`, `age`, `email`"));
assert!(sql.contains("'Bob', 25, 'bob@x.com'"));
}
#[test]
fn test_insert_empty_returns_empty() {
let sql = Query::insert().into_table("users").build();
assert_eq!(sql, "");
}
#[test]
fn test_insert_with_dialect() {
let sql = Query::insert()
.into_table("users")
.value("name", "'Alice'")
.build_with_dialect(DbType::PostgreSQL);
assert!(sql.contains("\"name\""));
assert!(sql.contains("\"users\""));
}
#[test]
fn test_update_basic() {
let sql = Query::update()
.table("users")
.set("name", "'Bob'")
.where_clause("id = 1")
.build();
assert!(sql.starts_with("UPDATE `users` SET"));
assert!(sql.contains("`name` = 'Bob'"));
assert!(sql.contains("WHERE id = 1"));
}
#[test]
fn test_update_multiple_sets() {
let sql = Query::update()
.table("users")
.sets(&[("name", "'Bob'"), ("age", "30")])
.where_clause("id = 1")
.build();
assert!(sql.contains("`name` = 'Bob', `age` = 30"));
}
#[test]
fn test_update_no_where() {
let sql = Query::update()
.table("users")
.set("status", "'active'")
.build();
assert!(sql.contains("UPDATE `users` SET `status` = 'active'"));
assert!(!sql.contains("WHERE"));
}
#[test]
fn test_update_empty_returns_empty() {
let sql = Query::update().table("users").build();
assert_eq!(sql, "");
}
#[test]
fn test_update_with_dialect() {
let sql = Query::update()
.table("users")
.set("name", "'Bob'")
.build_with_dialect(DbType::PostgreSQL);
assert!(sql.contains("\"users\""));
assert!(sql.contains("\"name\""));
}
#[test]
fn test_delete_basic() {
let sql = Query::delete()
.from_table("users")
.where_clause("id = 1")
.build();
assert!(sql.starts_with("DELETE FROM `users`"));
assert!(sql.contains("WHERE id = 1"));
}
#[test]
fn test_delete_no_where() {
let sql = Query::delete().from_table("users").build();
assert!(sql.contains("DELETE FROM `users`"));
assert!(!sql.contains("WHERE"));
}
#[test]
fn test_delete_multiple_wheres() {
let sql = Query::delete()
.from_table("users")
.where_clause("id > 100")
.where_clause("status = 'inactive'")
.build();
assert!(sql.contains("WHERE id > 100 AND status = 'inactive'"));
}
#[test]
fn test_delete_empty_returns_empty() {
let sql = Query::delete().build();
assert_eq!(sql, "");
}
#[test]
fn test_delete_with_dialect() {
let sql = Query::delete()
.from_table("users")
.where_clause("id = 1")
.build_with_dialect(DbType::PostgreSQL);
assert!(sql.contains("\"users\""));
}
#[test]
fn test_full_crud_flow() {
let insert = Query::insert()
.into_table("users")
.value("name", "'Alice'")
.value("age", "30")
.build();
assert!(insert.contains("INSERT INTO"));
let select = Query::select()
.column("id")
.column("name")
.from("users")
.where_clause("age > 18")
.order_by("id", true)
.limit(10)
.build(DbType::MySQL);
assert!(select.contains("SELECT"));
assert!(select.contains("FROM"));
assert!(select.contains("WHERE"));
assert!(select.contains("ORDER BY"));
assert!(select.contains("LIMIT"));
let update = Query::update()
.table("users")
.set("name", "'Bob'")
.where_clause("id = 1")
.build();
assert!(update.contains("UPDATE"));
assert!(update.contains("SET"));
assert!(update.contains("WHERE"));
let delete = Query::delete()
.from_table("users")
.where_clause("id = 1")
.build();
assert!(delete.contains("DELETE FROM"));
}
#[test]
fn test_complex_select_query() {
let sql = Query::select()
.distinct()
.columns(&["u.id", "u.name", "o.total"])
.from("users u")
.inner_join("orders o", "u.id = o.user_id")
.where_clause("u.status = 'active'")
.where_clause("o.total > 100")
.group_by("u.id")
.having("SUM(o.total) > 1000")
.order_by("u.id", true)
.limit(20)
.offset(40)
.build(DbType::MySQL);
assert!(sql.contains("SELECT DISTINCT"));
assert!(sql.contains("INNER JOIN `orders` o"));
assert!(sql.contains("WHERE u.status = 'active' AND o.total > 100"));
assert!(sql.contains("GROUP BY"));
assert!(sql.contains("HAVING SUM(o.total) > 1000"));
assert!(sql.contains("ORDER BY `u`.`id` ASC"));
assert!(sql.contains("LIMIT 20"));
assert!(sql.contains("OFFSET 40"));
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_select_where_rejects_semicolon_drop() {
let _ = Query::select()
.column("id")
.from("users")
.where_clause("1=1; DROP TABLE users")
.build(DbType::MySQL);
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_select_where_rejects_semicolon_space_drop() {
let _ = Query::select()
.column("id")
.from("users")
.where_clause("1=1; DROP TABLE users")
.build(DbType::MySQL);
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_select_where_rejects_line_comment() {
let _ = Query::select()
.column("id")
.from("users")
.where_clause("id = 1 -- DROP TABLE users")
.build(DbType::MySQL);
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_select_where_rejects_block_comment() {
let _ = Query::select()
.column("id")
.from("users")
.where_clause("id = 1 /* comment */ OR 1=1")
.build(DbType::MySQL);
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_select_or_where_rejects_drop() {
let _ = Query::select()
.column("id")
.from("users")
.where_clause("id = 1")
.or_where("1=1; DROP TABLE users")
.build(DbType::MySQL);
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_update_where_rejects_delete() {
let _ = Query::update()
.table("users")
.set("name", "'x'")
.where_clause("1=1; DELETE FROM users")
.build();
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_update_where_rejects_line_comment() {
let _ = Query::update()
.table("users")
.set("name", "'x'")
.where_clause("id = 1 -- bypass")
.build();
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_delete_where_rejects_drop() {
let _ = Query::delete()
.from_table("users")
.where_clause("1=1; DROP TABLE users")
.build();
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_delete_where_rejects_block_comment() {
let _ = Query::delete()
.from_table("users")
.where_clause("id = 1 /* */ OR 1=1")
.build();
}
#[test]
#[should_panic(expected = "SQL injection detected")]
fn test_delete_where_rejects_line_comment() {
let _ = Query::delete()
.from_table("users")
.where_clause("id = 1--")
.build();
}
#[test]
fn test_safe_where_clauses_pass() {
let _ = Query::select()
.column("id")
.from("users")
.where_clause("age > 18")
.where_clause("name = 'Alice;Bob'") .where_clause("id IN (1, 2, 3)")
.where_clause("created_at > '2026-01-01'")
.build(DbType::MySQL);
let _ = Query::update()
.table("users")
.set("name", "'x'")
.where_clause("id = 1")
.build();
let _ = Query::delete()
.from_table("users")
.where_clause("id = 1")
.build();
}
}