use sz_orm_core::{get_dialect, AggExpr, DbType, HavingOp, Model, QueryBuilder, Value};
#[derive(Debug, Clone, Default)]
struct Order {
id: i64,
}
impl Model for Order {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"orders"
}
fn pk(&self) -> Self::PrimaryKey {
self.id
}
fn set_pk(&mut self, pk: Self::PrimaryKey) {
self.id = pk;
}
}
fn mysql_builder() -> QueryBuilder<Order> {
QueryBuilder::<Order>::new(get_dialect(DbType::MySQL).unwrap())
}
#[test]
fn m5_having_invalid_aggregate_column_rejected() {
let qb = mysql_builder().table("orders").group_by("user_id");
let result = qb.having(
AggExpr::Sum("total`; DROP TABLE orders; --".to_string()),
HavingOp::Gt,
Value::I64(5),
);
assert!(
result.is_err(),
"注入列名必须被 having() 构建期拦截,实际为 Ok"
);
}
#[test]
fn m5_having_function_whitelist_only() {
let result = mysql_builder().having(
AggExpr::Sum("total".to_string()),
HavingOp::Gt,
Value::I64(5),
);
assert!(result.is_ok());
assert!(mysql_builder()
.having(AggExpr::CountStar, HavingOp::Gt, Value::I64(5))
.is_ok());
}
#[test]
fn m5_having_value_bound_as_param_not_inlined() {
let qb = mysql_builder()
.table("orders")
.group_by("user_id")
.having(
AggExpr::CountStar,
HavingOp::Gt,
Value::String("5 OR 1=1 --".to_string()),
)
.expect("valid aggregate");
let (sql, params) = qb.build_select_with_params();
assert!(
sql.contains("HAVING COUNT(*) > ?"),
"SQL 必须使用参数占位符,实际: {}",
sql
);
assert!(
!sql.contains("OR 1=1"),
"注入值不得内联进 SQL 文本,实际: {}",
sql
);
assert_eq!(params, vec![Value::String("5 OR 1=1 --".to_string())]);
}
#[test]
fn m5_having_valid_count_renders() {
let qb = mysql_builder()
.table("orders")
.group_by("user_id")
.having(AggExpr::CountStar, HavingOp::Gt, Value::I64(5))
.expect("valid aggregate");
let (sql, _params) = qb.build_select();
assert!(sql.contains("HAVING COUNT(*) > ?"), "实际: {}", sql);
let (sql, params) = qb.build_select_with_params();
assert!(sql.contains("HAVING COUNT(*) > ?"), "实际: {}", sql);
assert_eq!(params, vec![Value::I64(5)]);
}
#[test]
fn m5_having_sum_quoted_column() {
let qb = mysql_builder()
.table("orders")
.group_by("user_id")
.having(
AggExpr::Sum("total".to_string()),
HavingOp::Ge,
Value::I64(100),
)
.expect("valid aggregate");
let (sql, params) = qb.build_select_with_params();
assert!(sql.contains("HAVING SUM(`total`) >= ?"), "实际: {}", sql);
assert_eq!(params, vec![Value::I64(100)]);
}
#[test]
fn m5_having_multiple_conditions_and_joined() {
let qb = mysql_builder()
.table("orders")
.group_by("user_id")
.having(AggExpr::CountStar, HavingOp::Gt, Value::I64(5))
.expect("valid aggregate")
.having(
AggExpr::Sum("total".to_string()),
HavingOp::Lt,
Value::I64(1000),
)
.expect("valid aggregate");
let (sql, _params) = qb.build_select();
assert!(
sql.contains("HAVING COUNT(*) > ? AND SUM(`total`) < ?"),
"实际: {}",
sql
);
}
#[test]
fn m5_quick_query_having_parametized() {
use sz_orm_core::quick_query::Db;
let d = get_dialect(DbType::MySQL).unwrap();
let db = Db::new(d)
.name("orders")
.group_by("user_id")
.having(AggExpr::CountStar, HavingOp::Gt, Value::I64(5))
.expect("valid aggregate");
let (sql, _params) = db.build_select();
assert!(sql.contains("HAVING COUNT(*) > ?"), "实际: {}", sql);
}
#[test]
fn m6_select_invalid_column_rejected() {
let result = mysql_builder()
.table("users")
.select(vec!["id", "name`; DROP TABLE users; --"]);
assert!(
result.is_err(),
"注入列名必须被 select() 构建期拦截,实际为 Ok"
);
}
#[test]
fn m6_select_star_must_use_expr() {
assert!(mysql_builder().table("users").select(vec!["*"]).is_err());
let qb = mysql_builder().table("users").select_expr(vec!["*"]);
assert!(qb.build_select().0.contains("SELECT * FROM"));
}
#[test]
fn m6_select_valid_columns_quoted() {
let qb = mysql_builder()
.table("users")
.select(vec!["id", "name"])
.expect("valid columns");
let (sql, _params) = qb.build_select();
assert!(sql.contains("SELECT `id`, `name` FROM"), "实际: {}", sql);
}
#[test]
fn m6_select_expr_escape_hatch_raw() {
let qb = mysql_builder()
.table("users")
.select_expr(vec!["user_id", "COUNT(*) as cnt"]);
let (sql, _params) = qb.build_select();
assert!(
sql.contains("SELECT user_id, COUNT(*) as cnt FROM"),
"实际: {}",
sql
);
}
#[test]
fn m6_quick_query_select_validated() {
use sz_orm_core::quick_query::Db;
let d = get_dialect(DbType::MySQL).unwrap();
assert!(Db::new(d).name("users").select(vec!["id"]).is_ok());
let d2 = get_dialect(DbType::MySQL).unwrap();
assert!(Db::new(d2).name("users").select(vec!["id`; --"]).is_err());
}