use sea_query::{
Alias, Asterisk, Condition, Expr, ExprTrait, Func, LikeExpr, MysqlQueryBuilder, NullOrdering,
OnConflict, Order, OrderedStatement, PostgresQueryBuilder, Query, SelectStatement, SimpleExpr,
SqliteQueryBuilder, Value as SeaValue, WindowStatement,
};
use sea_query_sqlx::{SqlxBinder, SqlxValues};
use crate::config::QueryConfig;
use crate::query::IncludePlan;
use crate::query::backend::SqlDialect;
use crate::query::error::QueryError;
use crate::query::ir::{CmpOp, Cond, FieldRef, Quant, RelRef, TextOp, Value};
use crate::query::spec::{QuerySpec, SortKey};
use crate::query::write::{ConflictAction, ResolvedConflict, ResolvedWrite, WriteError};
pub const INCLUDE_RANK_COLUMN: &str = "__orion_include_rank";
const INCLUDE_SUBQUERY_ALIAS: &str = "__orion_include";
pub fn render(
spec: &QuerySpec,
cond: &Cond,
root_table: &str,
dialect: SqlDialect,
limits: &QueryConfig,
) -> Result<SelectStatement, QueryError> {
let limit = resolve_limit(spec.limit, limits)?;
let skip = resolve_skip(spec.skip, limits)?;
let mut stmt = Query::select();
match super::plan_projection(&spec.fields) {
None => {
stmt.column(Asterisk);
}
Some(fields) => {
for f in fields {
stmt.column(Alias::new(f.as_str()));
}
}
}
stmt.from(Alias::new(root_table));
if !matches!(cond, Cond::True) {
stmt.cond_where(render_expr(cond, root_table)?);
}
apply_sort_keys(&mut stmt, &spec.sort, dialect);
stmt.limit(limit);
if let Some(skip) = skip {
stmt.offset(skip);
}
Ok(stmt)
}
pub fn build_for(dialect: SqlDialect, stmt: &SelectStatement) -> (String, SqlxValues) {
match dialect {
SqlDialect::Postgres => stmt.build_sqlx(PostgresQueryBuilder),
SqlDialect::Mysql => stmt.build_sqlx(MysqlQueryBuilder),
SqlDialect::Sqlite => stmt.build_sqlx(SqliteQueryBuilder),
}
}
pub fn build_include_select(
inc: &IncludePlan,
keys: &[SeaValue],
dialect: SqlDialect,
) -> (String, SqlxValues) {
let foreign = inc.foreign.as_str();
let projection = inc.projection();
let mut inner = Query::select();
project_child(&mut inner, &projection);
let mut window = WindowStatement::partition_by(Alias::new(foreign));
apply_sort_keys(&mut window, &inc.sort, dialect);
inner.expr_window_as(
Func::cust(Alias::new("ROW_NUMBER")),
window,
Alias::new(INCLUDE_RANK_COLUMN),
);
inner.from(Alias::new(inc.target_table.as_str()));
inner.cond_where(Expr::col(Alias::new(foreign)).is_in(keys.to_vec()));
let mut stmt = Query::select();
project_child(&mut stmt, &projection);
stmt.from_subquery(inner, Alias::new(INCLUDE_SUBQUERY_ALIAS));
let cap = i64::try_from(inc.limit).unwrap_or(i64::MAX);
stmt.cond_where(Expr::col(Alias::new(INCLUDE_RANK_COLUMN)).lte(cap));
apply_sort_keys(&mut stmt, &inc.sort, dialect);
build_for(dialect, &stmt)
}
fn project_child(stmt: &mut SelectStatement, projection: &[String]) {
match super::plan_projection(projection) {
None => {
stmt.column(Asterisk);
}
Some(fields) => {
for f in fields {
stmt.column(Alias::new(f.as_str()));
}
}
}
}
pub fn json_key_to_sea(v: &serde_json::Value) -> Option<SeaValue> {
match v {
serde_json::Value::String(s) => Some(s.clone().into()),
serde_json::Value::Bool(b) => Some((*b).into()),
serde_json::Value::Number(n) => n
.as_i64()
.map(Into::into)
.or_else(|| n.as_f64().map(Into::into)),
_ => None,
}
}
use super::{resolve_limit, resolve_skip};
fn render_expr(cond: &Cond, current_table: &str) -> Result<SimpleExpr, QueryError> {
Ok(match cond {
Cond::True => Expr::val(1).eq(1),
Cond::False => Expr::val(1).eq(0),
Cond::And(cs) => fold_bool(cs, current_table, true)?,
Cond::Or(cs) => fold_bool(cs, current_table, false)?,
Cond::Not(inner) => render_expr(inner, current_table)?.not(),
Cond::Compare { field, op, value } => compare_expr(field, *op, value),
Cond::In {
field,
values,
negated,
} => in_expr(field, values, *negated),
Cond::IsNull { field, negated } => {
let col = col_expr(field);
if *negated {
col.is_not_null()
} else {
col.is_null()
}
}
Cond::Between {
field,
low,
high,
low_incl,
high_incl,
negated,
} => between_expr(field, low, high, *low_incl, *high_incl, *negated),
Cond::Text { field, op, pattern } => text_expr(field, *op, pattern),
Cond::Rel { quant, rel, cond } => render_rel(*quant, rel, cond, current_table)?,
})
}
fn fold_bool(cs: &[Cond], current_table: &str, and: bool) -> Result<SimpleExpr, QueryError> {
let mut iter = cs.iter();
let mut acc = match iter.next() {
Some(first) => render_expr(first, current_table)?,
None => return Ok(Expr::val(1).eq(if and { 1 } else { 0 })),
};
for c in iter {
let e = render_expr(c, current_table)?;
acc = if and { acc.and(e) } else { acc.or(e) };
}
Ok(acc)
}
fn render_rel(
quant: Quant,
rel: &RelRef,
inner: &Cond,
current_table: &str,
) -> Result<SimpleExpr, QueryError> {
let target = rel.target_table.as_str();
match quant {
Quant::Any => {
let inner_e = render_expr(inner, target)?;
Ok(Expr::exists(rel_subquery(
rel,
current_table,
Some(inner_e),
)?))
}
Quant::None => {
let inner_e = render_expr(inner, target)?;
Ok(Expr::exists(rel_subquery(rel, current_table, Some(inner_e))?).not())
}
Quant::All => {
let nonempty = Expr::exists(rel_subquery(rel, current_table, None)?);
let ie = render_expr(inner, target)?;
let violates = ie.clone().not().or(Expr::expr(ie).is_null());
let no_violation =
Expr::exists(rel_subquery(rel, current_table, Some(violates))?).not();
Ok(nonempty.and(no_violation))
}
}
}
fn rel_subquery(
rel: &RelRef,
current_table: &str,
inner: Option<SimpleExpr>,
) -> Result<SelectStatement, QueryError> {
let mut sub = Query::select();
sub.expr(Expr::val(1));
let mut where_cond = Condition::all();
match &rel.through {
None => {
sub.from(Alias::new(rel.target_table.as_str()));
where_cond = where_cond.add(
Expr::col((
Alias::new(rel.target_table.as_str()),
Alias::new(rel.foreign.as_str()),
))
.equals((Alias::new(current_table), Alias::new(rel.local.as_str()))),
);
}
Some(j) => {
sub.from(Alias::new(j.table.as_str()));
sub.inner_join(
Alias::new(rel.target_table.as_str()),
Expr::col((Alias::new(j.table.as_str()), Alias::new(j.foreign.as_str()))).equals((
Alias::new(rel.target_table.as_str()),
Alias::new(rel.foreign.as_str()),
)),
);
where_cond = where_cond.add(
Expr::col((Alias::new(j.table.as_str()), Alias::new(j.local.as_str())))
.equals((Alias::new(current_table), Alias::new(rel.local.as_str()))),
);
}
}
if let Some(e) = inner {
where_cond = where_cond.add(e);
}
sub.cond_where(where_cond);
Ok(sub)
}
fn compare_expr(field: &FieldRef, op: CmpOp, value: &Value) -> SimpleExpr {
let col = col_expr(field);
let v = to_sea_value(value);
match op {
CmpOp::Eq => col.eq(v),
CmpOp::Ne => col.ne(v),
CmpOp::Lt => col.lt(v),
CmpOp::Le => col.lte(v),
CmpOp::Gt => col.gt(v),
CmpOp::Ge => col.gte(v),
}
}
fn in_expr(field: &FieldRef, values: &[Value], negated: bool) -> SimpleExpr {
let col = col_expr(field);
let vals: Vec<SeaValue> = values.iter().map(to_sea_value).collect();
if negated {
col.is_not_in(vals)
} else {
col.is_in(vals)
}
}
fn between_expr(
field: &FieldRef,
low: &Value,
high: &Value,
low_incl: bool,
high_incl: bool,
negated: bool,
) -> SimpleExpr {
let lo = to_sea_value(low);
let hi = to_sea_value(high);
let e = if low_incl && high_incl {
col_expr(field).between(lo, hi)
} else {
let lo_e = if low_incl {
col_expr(field).gte(lo)
} else {
col_expr(field).gt(lo)
};
let hi_e = if high_incl {
col_expr(field).lte(hi)
} else {
col_expr(field).lt(hi)
};
lo_e.and(hi_e)
};
if negated { e.not() } else { e }
}
fn text_expr(field: &FieldRef, op: TextOp, pattern: &str) -> SimpleExpr {
let escaped = escape_like(pattern);
let like = match op {
TextOp::StartsWith => format!("{escaped}%"),
TextOp::EndsWith => format!("%{escaped}"),
TextOp::Contains => format!("%{escaped}%"),
};
col_expr(field).like(LikeExpr::new(like).escape('\\'))
}
fn escape_like(pattern: &str) -> String {
pattern
.replace('\\', "\\\\")
.replace('%', "\\%")
.replace('_', "\\_")
}
fn apply_sort_keys<S: OrderedStatement>(stmt: &mut S, sort: &[SortKey], dialect: SqlDialect) {
for p in super::plan_sort(sort) {
let order = if p.ascending { Order::Asc } else { Order::Desc };
let nulls = if p.nulls_first {
NullOrdering::First
} else {
NullOrdering::Last
};
match dialect {
SqlDialect::Mysql => {
stmt.order_by(Alias::new(p.field), order);
}
_ => {
stmt.order_by_with_nulls(Alias::new(p.field), order, nulls);
}
}
}
}
fn col_expr(field: &FieldRef) -> Expr {
Expr::col(Alias::new(field.physical.as_str()))
}
fn to_sea_value(v: &Value) -> SeaValue {
match v {
Value::Null => SeaValue::String(None),
Value::Bool(b) => (*b).into(),
Value::Int(i) => (*i).into(),
Value::Float(f) => (*f).into(),
Value::Str(s) => s.clone().into(),
}
}
pub fn render_write(
w: &ResolvedWrite,
dialect: SqlDialect,
) -> Result<(String, SqlxValues), WriteError> {
if !w.returning().is_empty() && dialect == SqlDialect::Mysql {
return Err(WriteError::Query(QueryError::FeatureUnsupportedByTarget {
feature: "returning".to_string(),
target: "mysql".to_string(),
}));
}
match w {
ResolvedWrite::Insert {
table,
columns,
rows,
returning,
} => render_insert(table, columns, rows, None, returning, dialect),
ResolvedWrite::Upsert {
table,
columns,
rows,
set,
conflict,
returning,
} => render_insert(
table,
columns,
rows,
Some((conflict, set)),
returning,
dialect,
),
ResolvedWrite::Update {
table,
set,
cond,
returning,
} => render_update(table, set, cond, returning, dialect),
ResolvedWrite::Delete {
table,
cond,
returning,
} => render_delete(table, cond, returning, dialect),
}
}
fn render_insert(
table: &str,
columns: &[String],
rows: &[Vec<Value>],
upsert: Option<(&ResolvedConflict, &[(String, Value)])>,
returning: &[String],
dialect: SqlDialect,
) -> Result<(String, SqlxValues), WriteError> {
let mut stmt = Query::insert();
stmt.into_table(Alias::new(table));
stmt.columns(columns.iter().map(|c| Alias::new(c.as_str())));
for row in rows {
let vals: Vec<SimpleExpr> = row.iter().map(value_expr).collect();
stmt.values(vals)
.map_err(|e| WriteError::Query(QueryError::InvalidEnvelope(e.to_string())))?;
}
if let Some((c, set)) = upsert {
let mut oc = OnConflict::columns(c.targets.iter().map(|t| Alias::new(t.as_str())));
match c.action {
ConflictAction::Nothing => {
oc.do_nothing();
}
ConflictAction::Update => {
if !set.is_empty() {
for (col, v) in set {
oc.value(Alias::new(col.as_str()), value_expr(v));
}
} else {
let upd: Vec<Alias> = columns
.iter()
.filter(|c2| !c.targets.contains(c2))
.map(|c2| Alias::new(c2.as_str()))
.collect();
if upd.is_empty() {
oc.do_nothing();
} else {
oc.update_columns(upd);
}
}
}
}
stmt.on_conflict(oc);
}
if !returning.is_empty() {
stmt.returning(
Query::returning().columns(returning.iter().map(|c| Alias::new(c.as_str()))),
);
}
Ok(build_write_for(dialect, &stmt))
}
fn render_update(
table: &str,
set: &[(String, Value)],
cond: &Option<Cond>,
returning: &[String],
dialect: SqlDialect,
) -> Result<(String, SqlxValues), WriteError> {
let mut stmt = Query::update();
stmt.table(Alias::new(table));
for (col, v) in set {
stmt.value(Alias::new(col.as_str()), value_expr(v));
}
if let Some(cond) = cond
&& !matches!(cond, Cond::True)
{
stmt.cond_where(render_expr(cond, table).map_err(WriteError::from)?);
}
if !returning.is_empty() {
stmt.returning(
Query::returning().columns(returning.iter().map(|c| Alias::new(c.as_str()))),
);
}
Ok(build_write_for(dialect, &stmt))
}
fn render_delete(
table: &str,
cond: &Option<Cond>,
returning: &[String],
dialect: SqlDialect,
) -> Result<(String, SqlxValues), WriteError> {
let mut stmt = Query::delete();
stmt.from_table(Alias::new(table));
if let Some(cond) = cond
&& !matches!(cond, Cond::True)
{
stmt.cond_where(render_expr(cond, table).map_err(WriteError::from)?);
}
if !returning.is_empty() {
stmt.returning(
Query::returning().columns(returning.iter().map(|c| Alias::new(c.as_str()))),
);
}
Ok(build_write_for(dialect, &stmt))
}
fn value_expr(v: &Value) -> SimpleExpr {
Expr::val(to_sea_value(v))
}
fn build_write_for<S: SqlxBinder>(dialect: SqlDialect, stmt: &S) -> (String, SqlxValues) {
match dialect {
SqlDialect::Postgres => stmt.build_sqlx(PostgresQueryBuilder),
SqlDialect::Mysql => stmt.build_sqlx(MysqlQueryBuilder),
SqlDialect::Sqlite => stmt.build_sqlx(SqliteQueryBuilder),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::query::{EntityRegistry, plan_sql};
use serde_json::{Value as Json, json};
fn limits() -> QueryConfig {
QueryConfig::default()
}
fn plan_stmt(
query: &Json,
reg: &EntityRegistry,
dialect: SqlDialect,
limits: &QueryConfig,
) -> Result<SelectStatement, QueryError> {
plan_sql(query, &serde_json::Map::new(), reg, dialect, limits).map(|plan| plan.main)
}
fn sql_for(query: Json, dialect: SqlDialect) -> String {
let stmt = plan_stmt(&query, &EntityRegistry::identity(), dialect, &limits())
.expect("translation should succeed");
match dialect {
SqlDialect::Sqlite => stmt.to_string(SqliteQueryBuilder),
SqlDialect::Postgres => stmt.to_string(PostgresQueryBuilder),
SqlDialect::Mysql => stmt.to_string(MysqlQueryBuilder),
}
}
fn sqlite(query: Json) -> String {
sql_for(query, SqlDialect::Sqlite)
}
fn rel_schema() -> EntityRegistry {
EntityRegistry::from_json(&json!({
"unmapped": "identity",
"entities": {
"users": {
"relations": {
"orders": { "to": "orders", "kind": "has_many", "local": "id", "foreign": "user_id" },
"tags": {
"to": "tags", "kind": "many_to_many", "local": "id", "foreign": "id",
"through": { "table": "user_tags", "local": "user_id", "foreign": "tag_id" }
}
}
}
}
}))
.expect("valid schema")
}
fn sqlite_schema(query: Json) -> String {
let stmt = plan_stmt(&query, &rel_schema(), SqlDialect::Sqlite, &limits())
.expect("translation should succeed");
stmt.to_string(SqliteQueryBuilder)
}
#[test]
fn test_select_all_no_filter() {
let sql = sqlite(json!({ "source": "users" }));
assert_eq!(sql, r#"SELECT * FROM "users" LIMIT 100"#);
}
#[test]
fn test_projection_and_comparison() {
let sql = sqlite(json!({
"source": "users",
"fields": ["id", "name"],
"filter": { ">": [{"field": "age"}, 18] }
}));
assert_eq!(
sql,
r#"SELECT "id", "name" FROM "users" WHERE "age" > 18 LIMIT 100"#
);
}
#[test]
fn test_and_or() {
let sql = sqlite(json!({
"source": "t",
"filter": { "and": [
{ "==": [{"field": "a"}, 1] },
{ "or": [ { "==": [{"field": "b"}, 2] }, { "==": [{"field": "c"}, 3] } ] }
] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" WHERE "a" = 1 AND ("b" = 2 OR "c" = 3) LIMIT 100"#
);
}
#[test]
fn test_membership() {
let sql = sqlite(json!({
"source": "t",
"filter": { "in": [{"field": "status"}, ["a", "b"]] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" WHERE "status" IN ('a', 'b') LIMIT 100"#
);
}
#[test]
fn test_empty_membership_is_false() {
let sql = sqlite(json!({
"source": "t",
"filter": { "in": [{"field": "status"}, []] }
}));
assert_eq!(sql, r#"SELECT * FROM "t" WHERE 1 = 0 LIMIT 100"#);
}
#[test]
fn test_is_null() {
let sql = sqlite(json!({
"source": "t",
"filter": { "==": [{"field": "email"}, null] }
}));
assert_eq!(sql, r#"SELECT * FROM "t" WHERE "email" IS NULL LIMIT 100"#);
}
#[test]
fn test_range_strict_is_not_between() {
let sql = sqlite(json!({
"source": "t",
"filter": { "<": [1, {"field": "x"}, 10] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" WHERE "x" > 1 AND "x" < 10 LIMIT 100"#
);
}
#[test]
fn test_range_inclusive_is_between() {
let sql = sqlite(json!({
"source": "t",
"filter": { "<=": [1, {"field": "x"}, 10] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" WHERE "x" BETWEEN 1 AND 10 LIMIT 100"#
);
}
#[test]
fn test_text_contains_escapes_wildcards() {
let sql = sqlite(json!({
"source": "t",
"filter": { "in": ["50%_off", {"field": "name"}] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" WHERE "name" LIKE '%50\%\_off%' ESCAPE '\' LIMIT 100"#
);
}
#[test]
fn test_starts_with() {
let sql = sqlite(json!({
"source": "t",
"filter": { "starts_with": [{"field": "name"}, "sm"] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" WHERE "name" LIKE 'sm%' ESCAPE '\' LIMIT 100"#
);
}
#[test]
fn test_sort_and_paging() {
let sql = sqlite(json!({
"source": "t",
"sort": [ { "created_at": "desc" } ],
"limit": 20,
"skip": 40
}));
assert_eq!(
sql,
r#"SELECT * FROM "t" ORDER BY "created_at" DESC NULLS LAST LIMIT 20 OFFSET 40"#
);
}
#[test]
fn test_null_ordering_is_nulls_smallest() {
assert_eq!(
sqlite(json!({ "source": "t", "sort": [ { "name": "asc" } ] })),
r#"SELECT * FROM "t" ORDER BY "name" ASC NULLS FIRST LIMIT 100"#
);
let stmt = plan_stmt(
&json!({ "source": "t", "sort": [ { "name": "asc" } ] }),
&EntityRegistry::identity(),
SqlDialect::Postgres,
&limits(),
)
.expect("ok");
assert_eq!(
stmt.to_string(PostgresQueryBuilder),
r#"SELECT * FROM "t" ORDER BY "name" ASC NULLS FIRST LIMIT 100"#
);
}
#[test]
fn test_postgres_placeholders_via_build() {
let stmt = plan_stmt(
&json!({ "source": "users", "filter": { "==": [{"field": "id"}, 7] } }),
&EntityRegistry::identity(),
SqlDialect::Postgres,
&limits(),
)
.expect("ok");
let (sql, _values) = build_for(SqlDialect::Postgres, &stmt);
assert_eq!(sql, r#"SELECT * FROM "users" WHERE "id" = $1 LIMIT $2"#);
}
#[test]
fn test_mysql_needs_no_null_ordering_emulation() {
let stmt = plan_stmt(
&json!({ "source": "t", "sort": [ { "name": "asc" } ] }),
&EntityRegistry::identity(),
SqlDialect::Mysql,
&limits(),
)
.expect("ok");
assert_eq!(
stmt.to_string(MysqlQueryBuilder),
"SELECT * FROM `t` ORDER BY `name` ASC LIMIT 100"
);
}
#[test]
fn test_limit_default_applied() {
let stmt = plan_stmt(
&json!({ "source": "t" }),
&EntityRegistry::identity(),
SqlDialect::Sqlite,
&QueryConfig {
default_limit: 50,
..QueryConfig::default()
},
)
.expect("ok");
assert_eq!(
stmt.to_string(SqliteQueryBuilder),
r#"SELECT * FROM "t" LIMIT 50"#
);
}
#[test]
fn test_limit_exceeds_max_rejected() {
let err = plan_stmt(
&json!({ "source": "t", "limit": 5000 }),
&EntityRegistry::identity(),
SqlDialect::Sqlite,
&limits(),
)
.expect_err("over the cap");
assert!(matches!(
err,
QueryError::LimitExceeded {
requested: 5000,
max: 1000
}
));
}
#[test]
fn test_skip_exceeds_max_rejected() {
let err = plan_stmt(
&json!({ "source": "t", "skip": 51 }),
&EntityRegistry::identity(),
SqlDialect::Sqlite,
&QueryConfig {
max_skip: 50,
..QueryConfig::default()
},
)
.expect_err("over the skip cap");
assert!(matches!(
err,
QueryError::SkipExceeded {
requested: 51,
max: 50
}
));
}
#[test]
fn test_relation_some_exists() {
let sql = sqlite_schema(json!({
"source": "users",
"filter": { "some": [{"field": "orders"}, {">": [{"field": "total"}, 100]}] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "users" WHERE EXISTS(SELECT 1 FROM "orders" WHERE "orders"."user_id" = "users"."id" AND "total" > 100) LIMIT 100"#
);
}
#[test]
fn test_relation_none_not_exists() {
let sql = sqlite_schema(json!({
"source": "users",
"filter": { "none": [{"field": "orders"}, {">": [{"field": "total"}, 100]}] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "users" WHERE NOT EXISTS(SELECT 1 FROM "orders" WHERE "orders"."user_id" = "users"."id" AND "total" > 100) LIMIT 100"#
);
}
#[test]
fn test_relation_all_null_fix() {
let sql = sqlite_schema(json!({
"source": "users",
"filter": { "all": [{"field": "orders"}, {">": [{"field": "total"}, 100]}] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "users" WHERE EXISTS(SELECT 1 FROM "orders" WHERE "orders"."user_id" = "users"."id") AND (NOT EXISTS(SELECT 1 FROM "orders" WHERE "orders"."user_id" = "users"."id" AND ((NOT "total" > 100) OR ("total" > 100) IS NULL))) LIMIT 100"#
);
}
#[test]
fn test_relation_many_to_many_join() {
let sql = sqlite_schema(json!({
"source": "users",
"filter": { "some": [{"field": "tags"}, {"==": [{"field": "label"}, "vip"]}] }
}));
assert_eq!(
sql,
r#"SELECT * FROM "users" WHERE EXISTS(SELECT 1 FROM "user_tags" INNER JOIN "tags" ON "user_tags"."tag_id" = "tags"."id" WHERE "user_tags"."user_id" = "users"."id" AND "label" = 'vip') LIMIT 100"#
);
}
#[test]
fn test_two_sibling_relations_are_independent() {
let sql = sqlite_schema(json!({
"source": "users",
"filter": { "and": [
{ "some": [{"field": "orders"}, {">": [{"field": "total"}, 100]}] },
{ "some": [{"field": "orders"}, {"==": [{"field": "user_id"}, 7]}] }
] }
}));
assert!(
sql.matches("EXISTS(SELECT 1 FROM \"orders\"").count() == 2,
"sql = {sql}"
);
}
fn plan_include(
selection: Json,
limits: &QueryConfig,
) -> Result<crate::query::SqlPlan, QueryError> {
crate::query::plan_sql(
&json!({
"source": "users",
"fields": ["name"],
"include": { "orders": selection }
}),
&serde_json::Map::new(),
&rel_schema(),
SqlDialect::Sqlite,
limits,
)
}
#[test]
fn test_plan_sql_augments_parent_key_and_plans_include() {
let plan = plan_include(
json!({ "fields": ["total"], "sort": [{ "id": "asc" }], "limit": 5 }),
&limits(),
)
.expect("plan");
assert_eq!(
plan.main.to_string(SqliteQueryBuilder),
r#"SELECT "name", "id" FROM "users" LIMIT 100"#
);
assert_eq!(plan.strip, vec!["id".to_string()]);
assert_eq!(plan.includes.len(), 1);
let inc = &plan.includes[0];
assert_eq!(inc.field, "orders");
assert_eq!(inc.target_table, "orders");
assert_eq!(inc.local, "id");
assert_eq!(inc.foreign, "user_id");
assert_eq!(inc.limit, 5);
assert_eq!(inc.sort.len(), 1);
}
#[test]
fn test_include_without_limit_takes_the_default_page_size() {
let plan = plan_include(json!({ "sort": [{ "id": "asc" }] }), &limits()).expect("plan");
assert_eq!(plan.includes[0].limit, 100);
}
#[test]
fn test_include_limit_over_the_cap_is_rejected() {
let err = plan_include(
json!({ "sort": [{ "id": "asc" }], "limit": 5000 }),
&limits(),
)
.expect_err("over the cap");
assert!(
matches!(
err,
QueryError::LimitExceeded {
requested: 5000,
max: 1000
}
),
"{err}"
);
}
#[test]
fn test_include_child_query_pages_per_parent_in_sql() {
let plan = plan_include(
json!({ "fields": ["total"], "sort": [{ "total": "desc" }], "limit": 2 }),
&limits(),
)
.expect("plan");
let keys = vec![SeaValue::from("u1"), SeaValue::from("u2")];
let (sql, _v) = build_include_select(&plan.includes[0], &keys, SqlDialect::Sqlite);
assert_eq!(
sql,
concat!(
r#"SELECT "total", "user_id" FROM (SELECT "total", "user_id", "#,
r#"ROW_NUMBER() OVER ( PARTITION BY "user_id" ORDER BY "total" DESC NULLS LAST ) "#,
r#"AS "__orion_include_rank" FROM "orders" WHERE "user_id" IN (?, ?)) "#,
r#"AS "__orion_include" WHERE "__orion_include_rank" <= ? "#,
r#"ORDER BY "total" DESC NULLS LAST"#
),
"sql = {sql}"
);
}
#[test]
fn test_include_projects_a_sort_key_it_was_not_asked_for() {
let plan = plan_include(
json!({ "fields": ["total"], "sort": [{ "created_at": "desc" }], "limit": 5 }),
&limits(),
)
.expect("plan");
let inc = &plan.includes[0];
assert_eq!(inc.projection(), ["total", "user_id", "created_at"]);
assert_eq!(inc.strip(), ["user_id", "created_at"]);
let keys = vec![SeaValue::from("u1")];
let (sql, _v) = build_include_select(inc, &keys, SqlDialect::Sqlite);
assert_eq!(
sql,
concat!(
r#"SELECT "total", "user_id", "created_at" FROM "#,
r#"(SELECT "total", "user_id", "created_at", "#,
r#"ROW_NUMBER() OVER ( PARTITION BY "user_id" ORDER BY "created_at" DESC NULLS LAST ) "#,
r#"AS "__orion_include_rank" FROM "orders" WHERE "user_id" IN (?)) "#,
r#"AS "__orion_include" WHERE "__orion_include_rank" <= ? "#,
r#"ORDER BY "created_at" DESC NULLS LAST"#
),
"sql = {sql}"
);
}
#[test]
fn test_include_projection_does_not_duplicate_or_over_strip() {
let plan = plan_include(
json!({ "fields": ["total", "user_id"], "sort": [{ "total": "asc" }] }),
&limits(),
)
.expect("plan");
let inc = &plan.includes[0];
assert_eq!(inc.projection(), ["total", "user_id"]);
assert!(inc.strip().is_empty(), "strip = {:?}", inc.strip());
}
#[test]
fn test_include_without_fields_projects_everything() {
let plan =
plan_include(json!({ "sort": [{ "created_at": "desc" }] }), &limits()).expect("plan");
let inc = &plan.includes[0];
assert!(inc.projection().is_empty());
assert!(inc.strip().is_empty());
let keys = vec![SeaValue::from("u1")];
let (sql, _v) = build_include_select(inc, &keys, SqlDialect::Sqlite);
assert!(
sql.starts_with(r#"SELECT * FROM (SELECT *, ROW_NUMBER() OVER ("#),
"sql = {sql}"
);
}
#[test]
fn test_include_child_query_renders_for_mysql() {
let plan = plan_include(
json!({ "fields": ["total"], "sort": [{ "total": "asc" }], "limit": 2 }),
&limits(),
)
.expect("plan");
let keys = vec![SeaValue::from("u1")];
let (sql, _v) = build_include_select(&plan.includes[0], &keys, SqlDialect::Mysql);
assert!(!sql.contains("NULLS"), "sql = {sql}");
assert!(
sql.contains("ROW_NUMBER() OVER ( PARTITION BY `user_id` ORDER BY `total` ASC )"),
"sql = {sql}"
);
}
#[test]
fn test_include_without_sort_is_rejected() {
let err = plan_include(json!({ "limit": 5 }), &limits()).expect_err("no order key");
assert!(matches!(err, QueryError::InvalidEnvelope(_)), "{err}");
assert!(err.to_string().contains("include.orders"), "{err}");
assert!(err.to_string().contains("sort"), "{err}");
}
#[test]
fn test_m2m_include_rejected() {
let err = crate::query::plan_sql(
&json!({ "source": "users", "include": { "tags": { "sort": [{ "id": "asc" }] } } }),
&serde_json::Map::new(),
&rel_schema(),
SqlDialect::Sqlite,
&limits(),
)
.expect_err("m2m include not supported");
assert!(matches!(err, QueryError::FeatureUnsupportedByTarget { .. }));
}
fn permissive_writes() -> crate::config::WriteConfig {
crate::config::WriteConfig {
max_rows: 1000,
allow_unfiltered: true,
}
}
fn write_sql(input: Json, dialect: SqlDialect) -> String {
let resolved = crate::query::write::resolve_write(
&input,
&serde_json::Map::new(),
&EntityRegistry::identity(),
&permissive_writes(),
)
.expect("resolve_write should succeed");
let (sql, _values) = render_write(&resolved, dialect).expect("render should succeed");
sql
}
fn sqlite_w(input: Json) -> String {
write_sql(input, SqlDialect::Sqlite)
}
#[test]
fn test_insert_single() {
let sql = sqlite_w(json!({
"op": "insert",
"target": "users",
"values": { "id": "u1", "name": "Alice" }
}));
assert_eq!(sql, r#"INSERT INTO "users" ("id", "name") VALUES (?, ?)"#);
}
#[test]
fn test_insert_bulk() {
let sql = sqlite_w(json!({
"op": "insert",
"target": "users",
"values": [ { "id": "u1", "name": "Alice" }, { "id": "u2", "name": "Bob" } ]
}));
assert_eq!(
sql,
r#"INSERT INTO "users" ("id", "name") VALUES (?, ?), (?, ?)"#
);
}
#[test]
fn test_insert_returning() {
let sql = sqlite_w(json!({
"op": "insert",
"target": "users",
"values": { "name": "Alice" },
"returning": ["id", "name"]
}));
assert_eq!(
sql,
r#"INSERT INTO "users" ("name") VALUES (?) RETURNING "id", "name""#
);
}
#[test]
fn test_update_with_filter() {
let sql = sqlite_w(json!({
"op": "update",
"target": "users",
"set": { "status": "inactive" },
"filter": { "==": [{ "field": "id" }, "u1"] }
}));
assert_eq!(sql, r#"UPDATE "users" SET "status" = ? WHERE "id" = ?"#);
}
#[test]
fn test_update_relation_filter_reuses_query_dialect() {
let input = json!({
"op": "update",
"target": "users",
"set": { "flagged": true },
"filter": { "some": [{ "field": "orders" }, { ">": [{ "field": "total" }, 100] }] }
});
let resolved = crate::query::write::resolve_write(
&input,
&serde_json::Map::new(),
&rel_schema(),
&permissive_writes(),
)
.expect("resolve");
let (sql, _v) = render_write(&resolved, SqlDialect::Sqlite).expect("render");
assert_eq!(
sql,
r#"UPDATE "users" SET "flagged" = ? WHERE EXISTS(SELECT ? FROM "orders" WHERE "orders"."user_id" = "users"."id" AND "total" > ?)"#
);
}
#[test]
fn test_delete_with_filter() {
let sql = sqlite_w(json!({
"op": "delete",
"target": "sessions",
"filter": { "<": [{ "field": "age" }, 0] }
}));
assert_eq!(sql, r#"DELETE FROM "sessions" WHERE "age" < ?"#);
}
#[test]
fn test_upsert_do_update() {
let sql = sqlite_w(json!({
"op": "upsert",
"target": "users",
"values": { "email": "a@x.io", "name": "Ada" },
"on_conflict": { "target": ["email"], "action": "update" }
}));
assert_eq!(
sql,
r#"INSERT INTO "users" ("email", "name") VALUES (?, ?) ON CONFLICT ("email") DO UPDATE SET "name" = "excluded"."name""#
);
}
#[test]
fn test_upsert_do_nothing() {
let sql = sqlite_w(json!({
"op": "upsert",
"target": "users",
"values": { "email": "a@x.io", "name": "Ada" },
"on_conflict": { "target": ["email"], "action": "nothing" }
}));
assert_eq!(
sql,
r#"INSERT INTO "users" ("email", "name") VALUES (?, ?) ON CONFLICT ("email") DO NOTHING"#
);
}
#[test]
fn test_returning_on_mysql_rejected() {
let input = json!({
"op": "insert",
"target": "users",
"values": { "name": "Ada" },
"returning": ["id"]
});
let resolved = crate::query::write::resolve_write(
&input,
&serde_json::Map::new(),
&EntityRegistry::identity(),
&permissive_writes(),
)
.expect("resolve");
let err = render_write(&resolved, SqlDialect::Mysql).expect_err("no RETURNING on MySQL");
assert!(matches!(
err,
crate::query::write::WriteError::Query(QueryError::FeatureUnsupportedByTarget { .. })
));
}
}
#[cfg(test)]
mod prop_tests {
use super::*;
use crate::config::WriteConfig;
use crate::query::schema::EntityRegistry;
use crate::query::write::resolve_write;
use proptest::prelude::*;
use serde_json::json;
fn permissive_writes() -> WriteConfig {
WriteConfig {
max_rows: 1000,
allow_unfiltered: true,
}
}
fn rendered(value: &str, dialect: SqlDialect) -> (Vec<String>, Vec<sea_query::Value>) {
let update = serde_json::json!({
"op": "update",
"target": "users",
"set": { "name": value },
"filter": { "==": [{"field": "name"}, value] }
});
let insert = serde_json::json!({
"op": "insert",
"target": "users",
"values": { "name": value }
});
let mut sqls = Vec::new();
let mut binds = Vec::new();
for input in [update, insert] {
let resolved = resolve_write(
&input,
&serde_json::Map::new(),
&EntityRegistry::identity(),
&permissive_writes(),
)
.expect("resolve");
let (sql, values) = render_write(&resolved, dialect).expect("render");
sqls.push(sql);
binds.extend(values.0.0);
}
(sqls, binds)
}
fn arb_ident() -> impl Strategy<Value = String> {
prop_oneof![
".{0,12}",
r#"[a-z]{0,3}["'`\\$.][a-z]{0,3}["'`\\$.]?"#,
"[A-Za-z_][A-Za-z0-9_]{0,8}",
]
}
fn is_hostile(ident: &str) -> bool {
ident.is_empty()
|| ident.starts_with('$')
|| ident.contains('.')
|| ident
.chars()
.any(|c| c.is_control() || matches!(c, '"' | '\'' | '`' | '\\'))
}
fn quoted(ident: &str, dialect: SqlDialect) -> String {
match dialect {
SqlDialect::Mysql => format!("`{ident}`"),
_ => format!("\"{ident}\""),
}
}
proptest! {
#[test]
fn sql_text_is_independent_of_user_values(value in ".*") {
for dialect in [SqlDialect::Postgres, SqlDialect::Sqlite, SqlDialect::Mysql] {
let (sqls, binds) = rendered(&value, dialect);
let (baseline_sqls, _) = rendered("baseline", dialect);
prop_assert_eq!(&sqls, &baseline_sqls, "dialect {:?}", dialect);
prop_assert!(
binds.iter().any(
|v| matches!(v, sea_query::Value::String(Some(s)) if s.as_str() == value)
),
"value must appear among the binds for {:?}",
dialect
);
}
}
#[test]
fn write_identifiers_are_rejected_or_safely_quoted(ident in arb_ident()) {
let mut values = serde_json::Map::new();
values.insert(ident.clone(), json!(1));
let insert = json!({
"op": "insert",
"target": ident.clone(),
"values": values.clone(),
});
let update = json!({
"op": "update",
"target": "users",
"set": values,
"returning": [ident.clone()],
"all": true,
});
for input in [insert, update] {
let resolved = resolve_write(
&input,
&serde_json::Map::new(),
&EntityRegistry::identity(),
&permissive_writes(),
);
match resolved {
Err(_) => {} Ok(w) => {
prop_assert!(
!is_hostile(&ident),
"resolve_write accepted a hostile identifier {ident:?}"
);
for dialect in [SqlDialect::Postgres, SqlDialect::Sqlite, SqlDialect::Mysql] {
let mut w = w.clone();
if dialect == SqlDialect::Mysql {
match &mut w {
ResolvedWrite::Insert { returning, .. }
| ResolvedWrite::Update { returning, .. }
| ResolvedWrite::Delete { returning, .. }
| ResolvedWrite::Upsert { returning, .. } => returning.clear(),
}
}
let (sql, _) = render_write(&w, dialect).expect("render");
prop_assert!(
sql.contains("ed(&ident, dialect)),
"identifier {ident:?} must render quoted in {sql:?}"
);
}
}
}
}
}
#[test]
fn query_identifiers_are_rejected_or_safely_quoted(ident in arb_ident()) {
let mut sort_key = serde_json::Map::new();
sort_key.insert(ident.clone(), json!("asc"));
let query = json!({
"source": ident.clone(),
"fields": [ident.clone()],
"sort": [sort_key],
"filter": { "==": [{ "field": ident.clone() }, 1] },
});
for dialect in [SqlDialect::Postgres, SqlDialect::Sqlite, SqlDialect::Mysql] {
match crate::query::plan_sql(
&query,
&serde_json::Map::new(),
&EntityRegistry::identity(),
dialect,
&QueryConfig::default(),
)
.map(|plan| plan.main)
{
Err(_) => {} Ok(stmt) => {
prop_assert!(
!is_hostile(&ident),
"plan_sql accepted a hostile identifier {ident:?}"
);
let sql = match dialect {
SqlDialect::Sqlite => stmt.to_string(SqliteQueryBuilder),
SqlDialect::Postgres => stmt.to_string(PostgresQueryBuilder),
SqlDialect::Mysql => stmt.to_string(MysqlQueryBuilder),
};
prop_assert!(
sql.contains("ed(&ident, dialect)),
"identifier {ident:?} must render quoted in {sql:?}"
);
}
}
}
}
}
}