use std::collections::HashMap;
use sz_orm_model::Dialect;
use sz_orm_model::Relation;
pub struct FindWithRelated<'a> {
dialect: &'a dyn Dialect,
main_table: String,
related_table: String,
foreign_key: String,
primary_key: String,
left_join: bool,
where_conds: Vec<String>,
order_by: Vec<(String, bool)>, limit: Option<usize>,
offset: Option<usize>,
}
impl<'a> FindWithRelated<'a> {
pub fn new(
dialect: &'a dyn Dialect,
main_table: impl Into<String>,
related_table: impl Into<String>,
foreign_key: impl Into<String>,
primary_key: impl Into<String>,
left_join: bool,
) -> Result<Self, sz_orm_model::DbError> {
let main_table = main_table.into();
let related_table = related_table.into();
let foreign_key = foreign_key.into();
let primary_key = primary_key.into();
validate_find_identifiers(&[&main_table, &related_table, &foreign_key, &primary_key])
.map_err(sz_orm_model::DbError::InvalidInput)?;
Ok(Self {
dialect,
main_table,
related_table,
foreign_key,
primary_key,
left_join,
where_conds: Vec::new(),
order_by: Vec::new(),
limit: None,
offset: None,
})
}
#[must_use]
pub fn where_cond(mut self, cond: impl Into<String>) -> Self {
self.where_conds.push(cond.into());
self
}
#[must_use]
pub fn order_by(mut self, field: impl Into<String>) -> Self {
self.order_by.push((field.into(), false));
self
}
#[must_use]
pub fn order_desc(mut self, field: impl Into<String>) -> Self {
self.order_by.push((field.into(), true));
self
}
#[must_use]
pub fn limit(mut self, n: usize) -> Self {
self.limit = Some(n);
self
}
#[must_use]
pub fn offset(mut self, n: usize) -> Self {
self.offset = Some(n);
self
}
pub fn build(&self) -> String {
let join_type = if self.left_join {
"LEFT JOIN"
} else {
"INNER JOIN"
};
let mut sql = format!(
"SELECT {}.*, {}.* FROM {} {} {} ON {}.{} = {}.{}",
self.dialect.quote(&self.main_table),
self.dialect.quote(&self.related_table),
self.dialect.quote(&self.main_table),
join_type,
self.dialect.quote(&self.related_table),
self.dialect.quote(&self.related_table),
self.dialect.quote(&self.foreign_key),
self.dialect.quote(&self.main_table),
self.dialect.quote(&self.primary_key),
);
if !self.where_conds.is_empty() {
sql.push_str(" WHERE ");
sql.push_str(&self.where_conds.join(" AND "));
}
if !self.order_by.is_empty() {
let parts: Vec<String> = self
.order_by
.iter()
.map(|(f, desc)| {
let d = if *desc { " DESC" } else { "" };
format!("{}{}", self.dialect.quote(f), d)
})
.collect();
sql.push_str(" ORDER BY ");
sql.push_str(&parts.join(", "));
}
if let Some(n) = self.limit {
sql.push_str(&format!(" LIMIT {}", n));
}
if let Some(n) = self.offset {
sql.push_str(&format!(" OFFSET {}", n));
}
sql
}
}
pub fn inspect_relation<'a>(
relations: &'a HashMap<&'a str, Relation>,
name: &'a str,
) -> Option<(&'a str, &'a str, &'a str, bool)> {
let rel = relations.get(name)?;
match rel {
Relation::HasMany(h) => Some((
h.child_model.as_str(),
h.foreign_key.as_str(),
h.child_pk.as_str(),
true,
)),
Relation::HasOne(h) => Some((
h.child_model.as_str(),
h.foreign_key.as_str(),
h.child_pk.as_str(),
false,
)),
Relation::BelongsTo(b) => Some((
b.parent_model.as_str(),
b.foreign_key.as_str(),
b.parent_pk.as_str(),
false,
)),
Relation::BelongsToMany(b) => Some((
b.target_model.as_str(),
b.foreign_key.as_str(),
b.other_key.as_str(),
true,
)),
Relation::MorphMany(m) => Some((
m.child_model.as_str(),
m.morph_id_column.as_str(),
"id",
true,
)),
Relation::MorphTo(m) => Some(("", m.morph_id_column.as_str(), "id", false)),
}
}
pub fn find_with_related_join<'a>(
dialect: &'a dyn Dialect,
main_table: &'a str,
related_table: &'a str,
foreign_key: &'a str,
primary_key: &'a str,
left_join: bool,
) -> Result<FindWithRelated<'a>, sz_orm_model::DbError> {
FindWithRelated::new(
dialect,
main_table,
related_table,
foreign_key,
primary_key,
left_join,
)
}
#[tracing::instrument(skip(dialect), fields(main_table = main_table, related_table = related_table, strategy = "eager_sql"))]
pub fn find_with_related_eager_sql(
dialect: &dyn Dialect,
main_table: &str,
related_table: &str,
foreign_key: &str,
main_where: Option<&str>,
) -> Result<(String, String), sz_orm_model::DbError> {
validate_find_identifiers(&[main_table, related_table, foreign_key])
.map_err(sz_orm_model::DbError::InvalidInput)?;
let main_sql = if let Some(w) = main_where {
format!("SELECT * FROM {} WHERE {}", dialect.quote(main_table), w)
} else {
format!("SELECT * FROM {}", dialect.quote(main_table))
};
let related_sql = format!(
"SELECT * FROM {} WHERE {} IN (?)",
dialect.quote(related_table),
dialect.quote(foreign_key),
);
Ok((main_sql, related_sql))
}
#[tracing::instrument(skip(dialect), fields(main_table = main_table, related_table = related_table, strategy = "subquery"))]
pub fn find_with_related_subquery(
dialect: &dyn Dialect,
main_table: &str,
related_table: &str,
foreign_key: &str,
primary_key: &str,
related_where: Option<&str>,
) -> Result<String, sz_orm_model::DbError> {
validate_find_identifiers(&[main_table, related_table, foreign_key, primary_key])
.map_err(sz_orm_model::DbError::InvalidInput)?;
let inner = if let Some(w) = related_where {
format!(
"SELECT {} FROM {} WHERE {}",
dialect.quote(foreign_key),
dialect.quote(related_table),
w
)
} else {
format!(
"SELECT {} FROM {}",
dialect.quote(foreign_key),
dialect.quote(related_table)
)
};
Ok(format!(
"SELECT * FROM {} WHERE {} IN ({})",
dialect.quote(main_table),
dialect.quote(primary_key),
inner
))
}
#[derive(Debug, Clone)]
enum WithRelationKind {
HasMany {
foreign_key: String,
primary_key: String,
},
HasOne {
foreign_key: String,
primary_key: String,
},
BelongsTo {
foreign_key: String,
primary_key: String,
},
}
#[derive(Debug, Clone)]
struct WithRelationItem {
related_table: String,
kind: WithRelationKind,
}
pub struct WithRelation<'a> {
dialect: &'a dyn Dialect,
main_table: String,
relations: Vec<(&'a str, WithRelationItem)>,
main_where: Option<String>,
}
impl<'a> WithRelation<'a> {
pub fn new(
dialect: &'a dyn Dialect,
main_table: impl Into<String>,
) -> Result<Self, sz_orm_model::DbError> {
let main_table = main_table.into();
validate_find_identifiers(&[&main_table]).map_err(sz_orm_model::DbError::InvalidInput)?;
Ok(Self {
dialect,
main_table,
relations: Vec::new(),
main_where: None,
})
}
pub fn with_has_many(
mut self,
related: &'a str,
foreign_key: impl Into<String>,
primary_key: impl Into<String>,
) -> Result<Self, sz_orm_model::DbError> {
let foreign_key = foreign_key.into();
let primary_key = primary_key.into();
validate_find_identifiers(&[related, &foreign_key, &primary_key])
.map_err(sz_orm_model::DbError::InvalidInput)?;
self.relations.push((
related,
WithRelationItem {
related_table: related.to_string(),
kind: WithRelationKind::HasMany {
foreign_key,
primary_key,
},
},
));
Ok(self)
}
pub fn with_has_one(
mut self,
related: &'a str,
foreign_key: impl Into<String>,
primary_key: impl Into<String>,
) -> Result<Self, sz_orm_model::DbError> {
let foreign_key = foreign_key.into();
let primary_key = primary_key.into();
validate_find_identifiers(&[related, &foreign_key, &primary_key])
.map_err(sz_orm_model::DbError::InvalidInput)?;
self.relations.push((
related,
WithRelationItem {
related_table: related.to_string(),
kind: WithRelationKind::HasOne {
foreign_key,
primary_key,
},
},
));
Ok(self)
}
pub fn with_belongs_to(
mut self,
related: &'a str,
foreign_key: impl Into<String>,
primary_key: impl Into<String>,
) -> Result<Self, sz_orm_model::DbError> {
let foreign_key = foreign_key.into();
let primary_key = primary_key.into();
validate_find_identifiers(&[related, &foreign_key, &primary_key])
.map_err(sz_orm_model::DbError::InvalidInput)?;
self.relations.push((
related,
WithRelationItem {
related_table: related.to_string(),
kind: WithRelationKind::BelongsTo {
foreign_key,
primary_key,
},
},
));
Ok(self)
}
#[tracing::instrument(skip(self), fields(strategy = "eager", main_table = &self.main_table))]
pub fn load_eager(mut self, main_where: Option<&str>) -> Result<Self, sz_orm_model::DbError> {
self.check_duplicate_relations()?;
self.main_where = main_where.map(String::from);
Ok(self)
}
fn check_duplicate_relations(&self) -> Result<(), sz_orm_model::DbError> {
let mut seen = std::collections::HashSet::new();
let mut duplicates = Vec::new();
for (name, _) in &self.relations {
if !seen.insert(*name) {
duplicates.push(*name);
}
}
if !duplicates.is_empty() {
return Err(sz_orm_model::DbError::InvalidInput(format!(
"WithRelation 重复关联检测失败:关联名 {:?} 被添加多次。请使用不同的关联名或移除重复项。",
duplicates
)));
}
Ok(())
}
#[tracing::instrument(skip(self), fields(strategy = "join", main_table = &self.main_table))]
pub fn load_join(&self, main_where: Option<&str>) -> Result<String, sz_orm_model::DbError> {
self.check_duplicate_relations()?;
let mut sql = format!("SELECT {}.*", self.dialect.quote(&self.main_table));
for (_, item) in &self.relations {
sql.push_str(&format!(", {}.*", self.dialect.quote(&item.related_table)));
}
sql.push_str(&format!(" FROM {}", self.dialect.quote(&self.main_table)));
for (_, item) in &self.relations {
let (join_type, left_col, right_col) = match &item.kind {
WithRelationKind::HasMany {
foreign_key,
primary_key,
}
| WithRelationKind::HasOne {
foreign_key,
primary_key,
} => (
"LEFT JOIN",
format!("{}.{}", item.related_table, foreign_key),
format!("{}.{}", self.main_table, primary_key),
),
WithRelationKind::BelongsTo {
foreign_key,
primary_key,
} => (
"INNER JOIN",
format!("{}.{}", self.main_table, foreign_key),
format!("{}.{}", item.related_table, primary_key),
),
};
let (l_table, l_col) = split_qualified(&left_col);
let (r_table, r_col) = split_qualified(&right_col);
sql.push_str(&format!(
" {} {} ON {}.{} = {}.{}",
join_type,
self.dialect.quote(&item.related_table),
self.dialect.quote(l_table),
self.dialect.quote(l_col),
self.dialect.quote(r_table),
self.dialect.quote(r_col),
));
}
if let Some(w) = main_where {
sql.push_str(&format!(" WHERE {}", w));
}
Ok(sql)
}
pub fn main_sql(&self) -> String {
let base = format!("SELECT * FROM {}", self.dialect.quote(&self.main_table));
if let Some(w) = &self.main_where {
format!("{} WHERE {}", base, w)
} else {
base
}
}
pub fn related_sql(&self, name: &str) -> Option<String> {
let (_, item) = self.relations.iter().find(|(n, _)| *n == name)?;
let foreign_key = match &item.kind {
WithRelationKind::HasMany { foreign_key, .. }
| WithRelationKind::HasOne { foreign_key, .. }
| WithRelationKind::BelongsTo { foreign_key, .. } => foreign_key.clone(),
};
Some(format!(
"SELECT * FROM {} WHERE {} IN (?)",
self.dialect.quote(&item.related_table),
self.dialect.quote(&foreign_key),
))
}
pub fn related_sql_with_ids(
&self,
name: &str,
ids: impl IntoIterator<Item = impl ToString>,
) -> Result<Option<String>, sz_orm_model::DbError> {
let (_, item) = self
.relations
.iter()
.find(|(n, _)| *n == name)
.ok_or_else(|| {
sz_orm_model::DbError::NotFound(format!("relation '{}' not found", name))
})?;
let foreign_key = match &item.kind {
WithRelationKind::HasMany { foreign_key, .. }
| WithRelationKind::HasOne { foreign_key, .. }
| WithRelationKind::BelongsTo { foreign_key, .. } => foreign_key.clone(),
};
let ids_str = ids
.into_iter()
.map(|v| {
let s = v.to_string();
sz_orm_model::sql_safety::validate_id_value(&s)?;
Ok(s)
})
.collect::<Result<Vec<_>, sz_orm_model::DbError>>()?
.join(", ");
Ok(Some(format!(
"SELECT * FROM {} WHERE {} IN ({})",
self.dialect.quote(&item.related_table),
self.dialect.quote(&foreign_key),
ids_str,
)))
}
pub fn relation_names(&self) -> Vec<&str> {
self.relations.iter().map(|(n, _)| *n).collect()
}
}
fn split_qualified(s: &str) -> (&str, &str) {
match s.rfind('.') {
Some(idx) => (&s[..idx], &s[idx + 1..]),
None => (s, ""),
}
}
fn is_valid_sql_identifier(s: &str) -> bool {
if s.is_empty() || s.len() > 64 {
return false;
}
let mut chars = s.chars();
match chars.next() {
Some(c) if c.is_ascii_alphabetic() || c == '_' => {}
_ => return false,
}
chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
}
fn validate_find_identifiers(idents: &[&str]) -> Result<(), String> {
for ident in idents {
if !is_valid_sql_identifier(ident) {
return Err(format!(
"invalid SQL identifier in find_with_related (potential SQL injection): {}",
ident
));
}
}
Ok(())
}
#[cfg(test)]
#[allow(deprecated)] mod tests {
use super::*;
use sz_orm_model::get_dialect;
use sz_orm_model::DbType;
use sz_orm_model::{BelongsTo, BelongsToMany, HasMany, HasOne, MorphMany, MorphTo};
fn mysql_dialect() -> Box<dyn Dialect> {
get_dialect(DbType::MySQL).expect("MySQL dialect")
}
fn pg_dialect() -> Box<dyn Dialect> {
get_dialect(DbType::PostgreSQL).expect("PG dialect")
}
fn sqlite_dialect() -> Box<dyn Dialect> {
get_dialect(DbType::Sqlite).expect("SQLite dialect")
}
#[test]
fn join_left_basic() {
let d = mysql_dialect();
let sql = FindWithRelated::new(&*d, "users", "profiles", "user_id", "id", true)
.unwrap()
.build();
assert!(sql.contains("SELECT `users`.*, `profiles`.*"));
assert!(sql.contains("FROM `users`"));
assert!(sql.contains("LEFT JOIN `profiles`"));
assert!(sql.contains("ON `profiles`.`user_id` = `users`.`id`"));
}
#[test]
fn join_inner_basic() {
let d = mysql_dialect();
let sql = FindWithRelated::new(&*d, "users", "orders", "user_id", "id", false)
.unwrap()
.build();
assert!(sql.contains("INNER JOIN `orders`"));
assert!(!sql.contains("LEFT JOIN"));
}
#[test]
fn join_with_where_order_limit() {
let d = mysql_dialect();
let sql = FindWithRelated::new(&*d, "users", "orders", "user_id", "id", true)
.unwrap()
.where_cond("users.status = 'active'")
.where_cond("orders.amount > 100")
.order_desc("orders.created_at")
.limit(10)
.offset(20)
.build();
assert!(sql.contains("WHERE users.status = 'active' AND orders.amount > 100"));
assert!(sql.contains("ORDER BY `orders.created_at` DESC"));
assert!(sql.contains("LIMIT 10"));
assert!(sql.contains("OFFSET 20"));
}
#[test]
fn join_pg_dialect() {
let d = pg_dialect();
let sql = FindWithRelated::new(&*d, "users", "orders", "user_id", "id", true)
.unwrap()
.build();
assert!(sql.contains("SELECT \"users\".*, \"orders\".*"));
assert!(sql.contains("LEFT JOIN \"orders\""));
assert!(sql.contains("ON \"orders\".\"user_id\" = \"users\".\"id\""));
}
#[test]
fn join_sqlite_dialect() {
let d = sqlite_dialect();
let sql = FindWithRelated::new(&*d, "users", "orders", "user_id", "id", true)
.unwrap()
.build();
assert!(sql.contains("LEFT JOIN \"orders\""));
}
#[test]
fn eager_sql_basic() {
let d = mysql_dialect();
let (main_sql, related_sql) =
find_with_related_eager_sql(&*d, "users", "orders", "user_id", Some("users.id > 0"))
.unwrap();
assert_eq!(main_sql, "SELECT * FROM `users` WHERE users.id > 0");
assert_eq!(related_sql, "SELECT * FROM `orders` WHERE `user_id` IN (?)");
}
#[test]
fn eager_sql_no_where() {
let d = mysql_dialect();
let (main_sql, related_sql) =
find_with_related_eager_sql(&*d, "users", "orders", "user_id", None).unwrap();
assert_eq!(main_sql, "SELECT * FROM `users`");
assert_eq!(related_sql, "SELECT * FROM `orders` WHERE `user_id` IN (?)");
}
#[test]
fn subquery_basic() {
let d = mysql_dialect();
let sql = find_with_related_subquery(
&*d,
"users",
"orders",
"user_id",
"id",
Some("orders.amount > 100"),
)
.unwrap();
assert_eq!(
sql,
"SELECT * FROM `users` WHERE `id` IN (SELECT `user_id` FROM `orders` WHERE orders.amount > 100)"
);
}
#[test]
fn subquery_no_where() {
let d = mysql_dialect();
let sql =
find_with_related_subquery(&*d, "users", "orders", "user_id", "id", None).unwrap();
assert_eq!(
sql,
"SELECT * FROM `users` WHERE `id` IN (SELECT `user_id` FROM `orders`)"
);
}
#[test]
fn inspect_relation_has_many() {
let mut rels = HashMap::new();
rels.insert(
"orders",
Relation::HasMany(HasMany {
foreign_key: "user_id".to_string(),
child_model: "orders".to_string(),
child_pk: "id".to_string(),
}),
);
let info = inspect_relation(&rels, "orders").expect("relation exists");
assert_eq!(info.0, "orders");
assert_eq!(info.1, "user_id");
assert_eq!(info.2, "id");
assert!(info.3, "HasMany 应标记为 is_many=true");
}
#[test]
fn inspect_relation_has_one() {
let mut rels = HashMap::new();
rels.insert(
"profile",
Relation::HasOne(HasOne {
foreign_key: "user_id".to_string(),
child_model: "profiles".to_string(),
child_pk: "id".to_string(),
}),
);
let info = inspect_relation(&rels, "profile").expect("relation exists");
assert_eq!(info.0, "profiles");
assert!(!info.3, "HasOne 应标记为 is_many=false");
}
#[test]
fn inspect_relation_belongs_to() {
let mut rels = HashMap::new();
rels.insert(
"user",
Relation::BelongsTo(BelongsTo {
foreign_key: "user_id".to_string(),
parent_model: "users".to_string(),
parent_pk: "id".to_string(),
}),
);
let info = inspect_relation(&rels, "user").expect("relation exists");
assert_eq!(info.0, "users");
assert!(!info.3, "BelongsTo 应标记为 is_many=false");
}
#[test]
fn inspect_relation_belongs_to_many() {
let mut rels = HashMap::new();
rels.insert(
"roles",
Relation::BelongsToMany(BelongsToMany {
junction_table: "user_role".to_string(),
foreign_key: "user_id".to_string(),
other_key: "role_id".to_string(),
target_model: "roles".to_string(),
target_pk: "id".to_string(),
}),
);
let info = inspect_relation(&rels, "roles").expect("relation exists");
assert_eq!(info.0, "roles");
assert!(info.3, "BelongsToMany 应标记为 is_many=true");
}
#[test]
fn inspect_relation_morph_many() {
let mut rels = HashMap::new();
rels.insert(
"comments",
Relation::MorphMany(MorphMany {
child_model: "comments".to_string(),
morph_type_column: "commentable_type".to_string(),
morph_id_column: "commentable_id".to_string(),
morph_type_value: "Post".to_string(),
}),
);
let info = inspect_relation(&rels, "comments").expect("relation exists");
assert_eq!(info.0, "comments");
assert_eq!(info.1, "commentable_id");
assert!(info.3, "MorphMany 应标记为 is_many=true");
}
#[test]
fn inspect_relation_morph_to() {
let mut rels = HashMap::new();
rels.insert(
"commentable",
Relation::MorphTo(MorphTo {
morph_type_column: "commentable_type".to_string(),
morph_id_column: "commentable_id".to_string(),
}),
);
let info = inspect_relation(&rels, "commentable").expect("relation exists");
assert_eq!(info.1, "commentable_id");
assert!(!info.3, "MorphTo 应标记为 is_many=false");
}
#[test]
fn inspect_relation_not_found() {
let rels = HashMap::<&str, Relation>::new();
assert!(inspect_relation(&rels, "nonexistent").is_none());
}
#[test]
fn empty_where_produces_no_where_clause() {
let d = mysql_dialect();
let sql = FindWithRelated::new(&*d, "users", "orders", "user_id", "id", true)
.unwrap()
.build();
assert!(!sql.contains("WHERE"));
}
#[test]
fn multiple_order_by() {
let d = mysql_dialect();
let sql = FindWithRelated::new(&*d, "users", "orders", "user_id", "id", true)
.unwrap()
.order_by("users.id")
.order_desc("orders.created_at")
.build();
assert!(sql.contains("ORDER BY `users.id`, `orders.created_at` DESC"));
}
#[test]
fn find_with_related_join_convenience_fn() {
let d = mysql_dialect();
let sql = find_with_related_join(&*d, "users", "orders", "user_id", "id", true)
.unwrap()
.where_cond("users.id = 1")
.build();
assert!(sql.contains("WHERE users.id = 1"));
}
#[test]
fn with_relation_load_has_many_eager() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(Some("users.id IN (1, 2, 3)"))
.expect("load_eager should succeed for non-duplicate relations");
assert_eq!(
loader.main_sql(),
"SELECT * FROM `users` WHERE users.id IN (1, 2, 3)"
);
assert_eq!(
loader.related_sql("orders").unwrap(),
"SELECT * FROM `orders` WHERE `user_id` IN (?)"
);
assert_eq!(
loader
.related_sql_with_ids("orders", [1_i64, 2, 3])
.unwrap()
.unwrap(),
"SELECT * FROM `orders` WHERE `user_id` IN (1, 2, 3)"
);
assert_eq!(
loader
.related_sql_with_ids("orders", ["a", "b", "c"])
.unwrap()
.unwrap(),
"SELECT * FROM `orders` WHERE `user_id` IN (a, b, c)"
);
}
#[test]
fn with_relation_load_has_one_join() {
let d = mysql_dialect();
let sql = WithRelation::new(&*d, "users")
.unwrap()
.with_has_one("profiles", "user_id", "id")
.unwrap()
.load_join(Some("users.id = 1"))
.expect("load_join should succeed for non-duplicate relations");
assert!(sql.contains("LEFT JOIN `profiles`"));
assert!(sql.contains("ON `profiles`.`user_id` = `users`.`id`"));
assert!(sql.contains("WHERE users.id = 1"));
}
#[test]
fn with_relation_load_belongs_to_join() {
let d = mysql_dialect();
let sql = WithRelation::new(&*d, "orders")
.unwrap()
.with_belongs_to("users", "user_id", "id")
.unwrap()
.load_join(None)
.expect("load_join should succeed for non-duplicate relations");
assert!(sql.contains("INNER JOIN `users`"));
assert!(sql.contains("ON `orders`.`user_id` = `users`.`id`"));
}
#[test]
fn with_relation_multiple_relations_eager() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.with_has_many("posts", "author_id", "id")
.unwrap()
.with_has_one("profiles", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
assert_eq!(loader.main_sql(), "SELECT * FROM `users`");
assert!(loader.related_sql("orders").is_some());
assert!(loader.related_sql("posts").is_some());
assert!(loader.related_sql("profiles").is_some());
assert!(loader.related_sql("nonexistent").is_none());
}
#[test]
fn with_relation_load_eager_with_specific_ids() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
let orders_sql = loader
.related_sql_with_ids("orders", [1_i64, 5, 10])
.unwrap()
.unwrap();
assert_eq!(
orders_sql,
"SELECT * FROM `orders` WHERE `user_id` IN (1, 5, 10)"
);
}
#[test]
fn with_relation_pg_dialect_eager() {
let d = pg_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(Some("users.id > 100"))
.expect("load_eager should succeed for non-duplicate relations");
assert_eq!(
loader.main_sql(),
"SELECT * FROM \"users\" WHERE users.id > 100"
);
assert_eq!(
loader.related_sql("orders").unwrap(),
"SELECT * FROM \"orders\" WHERE \"user_id\" IN (?)"
);
assert_eq!(
loader
.related_sql_with_ids("orders", [100_i64])
.unwrap()
.unwrap(),
"SELECT * FROM \"orders\" WHERE \"user_id\" IN (100)"
);
}
#[test]
fn with_relation_sqlite_dialect_join() {
let d = sqlite_dialect();
let sql = WithRelation::new(&*d, "users")
.unwrap()
.with_has_one("profiles", "user_id", "id")
.unwrap()
.load_join(None)
.expect("load_join should succeed for non-duplicate relations");
assert!(sql.contains("LEFT JOIN \"profiles\""));
}
#[test]
fn with_relation_rejects_sql_injection_in_id_semicolon() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
let err = loader
.related_sql_with_ids("orders", ["1; DROP TABLE users"])
.unwrap_err();
assert!(matches!(err, sz_orm_model::DbError::InvalidInput(_)));
}
#[test]
fn with_relation_rejects_sql_injection_in_id_or() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
let err = loader
.related_sql_with_ids("orders", ["1) OR 1=1"])
.unwrap_err();
assert!(matches!(err, sz_orm_model::DbError::InvalidInput(_)));
}
#[test]
fn with_relation_rejects_sql_injection_in_id_quote() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
let err = loader
.related_sql_with_ids("orders", ["' OR '1'='1"])
.unwrap_err();
assert!(matches!(err, sz_orm_model::DbError::InvalidInput(_)));
}
#[test]
fn with_relation_rejects_sql_injection_in_id_comment() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
let err = loader.related_sql_with_ids("orders", ["1--"]).unwrap_err();
assert!(matches!(err, sz_orm_model::DbError::InvalidInput(_)));
}
#[test]
fn with_relation_rejects_sql_injection_in_id_with_space() {
let d = mysql_dialect();
let loader = WithRelation::new(&*d, "users")
.unwrap()
.with_has_many("orders", "user_id", "id")
.unwrap()
.load_eager(None)
.expect("load_eager should succeed for non-duplicate relations");
let err = loader.related_sql_with_ids("orders", ["1 2"]).unwrap_err();
assert!(matches!(err, sz_orm_model::DbError::InvalidInput(_)));
}
}