#![cfg(test)]
use std::collections::HashMap;
use sz_orm_core::dialect::get_dialect;
use sz_orm_core::{DbType, Model, ModelExt, QueryBuilder, Relation, Value};
#[derive(Clone, Debug)]
#[allow(dead_code)]
struct User {
id: i64,
name: String,
age: i32,
email: String,
}
impl Model for User {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"users"
}
fn pk(&self) -> Self::PrimaryKey {
self.id
}
fn set_pk(&mut self, pk: Self::PrimaryKey) {
self.id = pk;
}
}
impl ModelExt for User {
fn columns() -> Vec<&'static str> {
vec!["id", "name", "age", "email"]
}
fn fillable() -> Vec<&'static str> {
vec!["name", "age", "email"]
}
fn guarded() -> Vec<&'static str> {
vec!["id"]
}
fn hidden() -> Vec<&'static str> {
vec![]
}
fn relations() -> HashMap<&'static str, Relation> {
HashMap::new()
}
fn fill(&mut self, _data: HashMap<String, Value>) {}
fn to_json(&self) -> serde_json::Value {
serde_json::json!({})
}
}
fn make_builder(db_type: DbType) -> QueryBuilder<User> {
let dialect = get_dialect(db_type).unwrap();
QueryBuilder::<User>::new(dialect).table("users")
}
fn make_row(id: i64, name: &str, age: i32, email: &str) -> HashMap<String, Value> {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(id));
row.insert("name".to_string(), Value::String(name.to_string()));
row.insert("age".to_string(), Value::I32(age));
row.insert("email".to_string(), Value::String(email.to_string()));
row
}
fn strip_quotes(sql: &str) -> String {
sql.replace(['`', '"'], "")
}
#[test]
fn test_l3_1_mysql_batch_insert_basic() {
let builder = make_builder(DbType::MySQL);
let rows = vec![
make_row(1, "Alice", 30, "alice@test.com"),
make_row(2, "Bob", 25, "bob@test.com"),
];
let (sql, params) = builder.build_batch_insert_with_params(&rows);
assert!(
sql.to_uppercase().starts_with("INSERT INTO"),
"应以 INSERT INTO 开头: {}",
sql
);
assert!(
sql.to_uppercase().contains("VALUES"),
"应包含 VALUES: {}",
sql
);
let values_pos = sql.to_uppercase().find("VALUES").unwrap();
let values_part = &sql[values_pos..];
let paren_count = values_part.matches('(').count();
assert_eq!(paren_count, 2, "VALUES 后应有 2 组值括号: {}", sql);
assert_eq!(params.len(), 8, "应有 8 个参数: {:?}", params);
}
#[test]
fn test_l3_2_mysql_batch_insert_param_order() {
let builder = make_builder(DbType::MySQL);
let rows = vec![make_row(10, "Alice", 30, "a@t.com")];
let (sql, params) = builder.build_batch_insert_with_params(&rows);
assert_eq!(params.len(), 4, "单行应有 4 个参数: {:?}", params);
let placeholder_count = sql.matches('?').count();
assert_eq!(placeholder_count, 4, "应有 4 个 ? 占位符: {}", sql);
}
#[test]
fn test_l3_3_mysql_batch_upsert_on_duplicate_key() {
let builder = make_builder(DbType::MySQL);
let rows = vec![
make_row(1, "Alice", 30, "alice@t.com"),
make_row(2, "Bob", 25, "bob@t.com"),
];
let result = builder.build_batch_upsert_with_params(&rows, &["id"], &[]);
assert!(result.is_ok(), "MySQL upsert 应成功: {:?}", result);
let (sql, params) = result.unwrap();
let sql_upper = sql.to_uppercase();
assert!(
sql_upper.contains("ON DUPLICATE KEY UPDATE"),
"应包含 ON DUPLICATE KEY UPDATE: {}",
sql
);
assert!(
sql_upper.contains("VALUES("),
"应包含 VALUES() 函数引用: {}",
sql
);
assert_eq!(params.len(), 8, "应有 8 个参数: {:?}", params);
}
#[test]
fn test_l3_4_mysql_upsert_specific_update_columns() {
let builder = make_builder(DbType::MySQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let (sql, _) = builder
.build_batch_upsert_with_params(&rows, &["id"], &["name", "age"])
.unwrap();
let sql_clean = strip_quotes(&sql);
assert!(
sql_clean.contains("name=VALUES(name)"),
"应更新 name 列: {}",
sql
);
assert!(
sql_clean.contains("age=VALUES(age)"),
"应更新 age 列: {}",
sql
);
assert!(
!sql_clean.contains("email=VALUES(email)"),
"不应更新 email 列(未指定): {}",
sql
);
}
#[test]
fn test_l3_5_mysql_upsert_update_all_columns() {
let builder = make_builder(DbType::MySQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let (sql, _) = builder
.build_batch_upsert_with_params(&rows, &["id"], &[])
.unwrap();
let sql_clean = strip_quotes(&sql);
assert!(
sql_clean.contains("name=VALUES(name)"),
"应更新 name: {}",
sql
);
assert!(sql_clean.contains("age=VALUES(age)"), "应更新 age: {}", sql);
assert!(
sql_clean.contains("email=VALUES(email)"),
"应更新 email: {}",
sql
);
}
#[test]
fn test_l3_6_pg_batch_upsert_on_conflict() {
let builder = make_builder(DbType::PostgreSQL);
let rows = vec![
make_row(1, "Alice", 30, "alice@t.com"),
make_row(2, "Bob", 25, "bob@t.com"),
];
let result = builder.build_batch_upsert_with_params(&rows, &["id"], &[]);
assert!(result.is_ok(), "PG upsert 应成功: {:?}", result);
let (sql, params) = result.unwrap();
let sql_upper = sql.to_uppercase();
assert!(
sql_upper.contains("ON CONFLICT"),
"应包含 ON CONFLICT: {}",
sql
);
assert!(
sql_upper.contains("DO UPDATE SET"),
"应包含 DO UPDATE SET: {}",
sql
);
assert!(
sql_upper.contains("EXCLUDED"),
"应包含 EXCLUDED 引用: {}",
sql
);
assert_eq!(params.len(), 8, "应有 8 个参数: {:?}", params);
}
#[test]
fn test_l3_7_pg_upsert_conflict_columns_in_clause() {
let builder = make_builder(DbType::PostgreSQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let (sql, _) = builder
.build_batch_upsert_with_params(&rows, &["id"], &[])
.unwrap();
let sql_clean = strip_quotes(&sql);
assert!(
sql_clean.to_uppercase().contains("ON CONFLICT (ID)"),
"ON CONFLICT 应包含 id 列: {}",
sql
);
}
#[test]
fn test_l3_8_pg_upsert_update_non_conflict_columns() {
let builder = make_builder(DbType::PostgreSQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let (sql, _) = builder
.build_batch_upsert_with_params(&rows, &["id"], &[])
.unwrap();
let sql_clean = strip_quotes(&sql);
assert!(
sql_clean.contains("name=EXCLUDED.name"),
"应更新 name: {}",
sql
);
assert!(
sql_clean.contains("age=EXCLUDED.age"),
"应更新 age: {}",
sql
);
assert!(
sql_clean.contains("email=EXCLUDED.email"),
"应更新 email: {}",
sql
);
assert!(
!sql_clean.contains("id=EXCLUDED.id"),
"不应更新冲突列 id: {}",
sql
);
}
#[test]
fn test_l3_9_pg_upsert_specific_update_columns() {
let builder = make_builder(DbType::PostgreSQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let (sql, _) = builder
.build_batch_upsert_with_params(&rows, &["id"], &["name"])
.unwrap();
let sql_clean = strip_quotes(&sql);
assert!(
sql_clean.contains("name=EXCLUDED.name"),
"应更新 name: {}",
sql
);
assert!(
!sql_clean.contains("age=EXCLUDED.age"),
"不应更新 age(未指定): {}",
sql
);
assert!(
!sql_clean.contains("email=EXCLUDED.email"),
"不应更新 email(未指定): {}",
sql
);
}
#[test]
fn test_l3_10_sqlite_batch_upsert() {
let builder = make_builder(DbType::Sqlite);
let rows = vec![
make_row(1, "Alice", 30, "alice@t.com"),
make_row(2, "Bob", 25, "bob@t.com"),
];
let result = builder.build_batch_upsert_with_params(&rows, &["id"], &[]);
assert!(result.is_ok(), "SQLite upsert 应成功: {:?}", result);
let (sql, params) = result.unwrap();
let sql_upper = sql.to_uppercase();
assert!(
sql_upper.contains("ON CONFLICT"),
"应包含 ON CONFLICT: {}",
sql
);
assert!(
sql_upper.contains("DO UPDATE SET"),
"应包含 DO UPDATE SET: {}",
sql
);
assert!(sql_upper.contains("EXCLUDED"), "应包含 EXCLUDED: {}", sql);
assert_eq!(params.len(), 8, "应有 8 个参数: {:?}", params);
}
#[test]
fn test_l3_11_empty_rows_returns_error() {
let builder = make_builder(DbType::MySQL);
let rows: Vec<HashMap<String, Value>> = vec![];
let result = builder.build_batch_upsert_with_params(&rows, &["id"], &[]);
assert!(result.is_err(), "空行应返回错误");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("empty") || err_msg.contains("不能为空"),
"错误消息应说明空行问题: {}",
err_msg
);
}
#[test]
fn test_l3_12_single_row_upsert() {
let builder = make_builder(DbType::MySQL);
let rows = vec![make_row(42, "Single", 99, "single@t.com")];
let (sql, params) = builder
.build_batch_upsert_with_params(&rows, &["id"], &[])
.unwrap();
assert!(
sql.to_uppercase().contains("ON DUPLICATE KEY UPDATE"),
"应包含 upsert 子句: {}",
sql
);
assert_eq!(params.len(), 4, "单行应有 4 个参数: {:?}", params);
}
#[test]
fn test_l3_13_oracle_unsupported_upsert() {
let builder = make_builder(DbType::Oracle);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let result = builder.build_batch_upsert_with_params(&rows, &["id"], &[]);
assert!(result.is_err(), "Oracle 应返回不支持错误");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("does not support") || err_msg.contains("不支持"),
"错误消息应说明不支持: {}",
err_msg
);
}
#[test]
fn test_l3_14_sqlserver_unsupported_upsert() {
let builder = make_builder(DbType::SqlServer);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let result = builder.build_batch_upsert_with_params(&rows, &["id"], &[]);
assert!(result.is_err(), "SQL Server 应返回不支持错误");
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("does not support") || err_msg.contains("不支持"),
"错误消息应说明不支持: {}",
err_msg
);
assert!(
err_msg.contains("upsert"),
"错误消息应包含 upsert 关键词: {}",
err_msg
);
assert!(
err_msg.contains("SqlServer") || err_msg.contains("SQL Server"),
"错误消息应指明 SqlServer 方言: {}",
err_msg
);
}
#[test]
fn test_l3_15_pg_no_conflict_columns_error() {
let builder = make_builder(DbType::PostgreSQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let result = builder.build_batch_upsert_with_params(&rows, &[], &[]);
assert!(result.is_err(), "PG 无冲突列应返回错误");
}
#[test]
fn test_l3_16_mysql_no_conflict_columns_ok() {
let builder = make_builder(DbType::MySQL);
let rows = vec![make_row(1, "Alice", 30, "alice@t.com")];
let result = builder.build_batch_upsert_with_params(&rows, &[], &[]);
assert!(
result.is_ok(),
"MySQL 无冲突列应成功(自动检测唯一键): {:?}",
result
);
let (sql, _) = result.unwrap();
assert!(
sql.to_uppercase().contains("ON DUPLICATE KEY UPDATE"),
"应包含 upsert: {}",
sql
);
}
#[test]
fn test_l3_17_null_value_handling() {
let builder = make_builder(DbType::MySQL);
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::Null);
row.insert("age".to_string(), Value::I32(30));
row.insert("email".to_string(), Value::String("a@t.com".to_string()));
let (sql, params) = builder.build_batch_insert_with_params(&[row]);
assert_eq!(
params.len(),
4,
"应有 4 个参数(Null 也占参数): {:?}",
params
);
let placeholder_count = sql.matches('?').count();
assert_eq!(placeholder_count, 4, "应有 4 个 ? 占位符: {}", sql);
}
#[test]
fn test_l3_18_unicode_values() {
let builder = make_builder(DbType::MySQL);
let rows = vec![make_row(1, "张三", 25, "张三@测试.com")];
let (sql, params) = builder
.build_batch_upsert_with_params(&rows, &["id"], &[])
.unwrap();
assert!(
!sql.contains("张三"),
"中文值不应出现在 SQL 中(应参数化): {}",
sql
);
assert_eq!(params.len(), 4, "应有 4 个参数: {:?}", params);
}
#[test]
fn test_l3_19_sql_injection_prevention() {
let builder = make_builder(DbType::MySQL);
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert(
"name".to_string(),
Value::String("'; DROP TABLE users; --".to_string()),
);
row.insert("age".to_string(), Value::I32(30));
row.insert("email".to_string(), Value::String("safe@t.com".to_string()));
let (sql, params) = builder
.build_batch_upsert_with_params(&[row], &["id"], &[])
.unwrap();
assert!(
!sql.contains("DROP TABLE"),
"SQL 注入字符串不应出现在 SQL 中: {}",
sql
);
assert!(!sql.contains("--"), "SQL 注释不应出现在 SQL 中: {}", sql);
assert_eq!(params.len(), 4, "应有 4 个参数: {:?}", params);
}
#[test]
fn test_l3_20_large_batch_100_rows() {
let builder = make_builder(DbType::PostgreSQL);
let rows: Vec<HashMap<String, Value>> = (1..=100)
.map(|i| {
make_row(
i,
&format!("user{}", i),
(i % 80) as i32,
&format!("u{}@t.com", i),
)
})
.collect();
let (sql, params) = builder
.build_batch_upsert_with_params(&rows, &["id"], &[])
.unwrap();
assert_eq!(
params.len(),
400,
"100 行应有 400 个参数: {:?}",
params.len()
);
let placeholder_count = sql.matches('?').count();
assert_eq!(
placeholder_count, 400,
"应有 400 个 ? 占位符: {}",
placeholder_count
);
assert!(
sql.to_uppercase().contains("ON CONFLICT"),
"应包含 ON CONFLICT: {}",
sql
);
}