use crate::model::Model;
use crate::pg::accumulator::SqlAccumulator;
use crate::pg::decode::{FromPgRow, joined_alias_for_prefix};
use crate::query::condition::{Condition, FilterValue, Leaf, LookupOp};
use crate::query::portable::{PortablePredicateError, SqlEmitContext};
use crate::query::q::{ArrayPredicate, CompoundOp, Q};
use crate::query::queryset::{DistinctMode, QuerySet};
use crate::query::merge::{
MergeAction, MergeBranch, MergeMatchKind, MergeOnEq, MergeValue, SRC_ALIAS, TGT_ALIAS,
};
fn escape_like(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'\\' | '%' | '_' => {
out.push('\\');
out.push(c);
}
_ => out.push(c),
}
}
out
}
pub(crate) fn push_filter_value(acc: &mut SqlAccumulator, v: FilterValue) {
match v {
FilterValue::String(s) => {
acc.push_bind(s);
}
FilterValue::I16(n) => {
acc.push_bind(n);
}
FilterValue::I32(n) => {
acc.push_bind(n);
}
FilterValue::I64(n) => {
acc.push_bind(n);
}
FilterValue::F32(n) => {
acc.push_bind(n);
}
FilterValue::F64(n) => {
acc.push_bind(n);
}
FilterValue::Bool(b) => {
acc.push_bind(b);
}
FilterValue::Timestamp(d) => {
acc.push_bind(d);
}
FilterValue::DateTime(d) => {
acc.push_bind(d);
}
FilterValue::Date(d) => {
acc.push_bind(d);
}
FilterValue::Uuid(u) => {
acc.push_bind(u);
}
FilterValue::HeerId(h) => {
acc.push_bind(h);
}
FilterValue::RanjId(r) => {
acc.push_bind(r);
}
FilterValue::HeerIdDesc(h) => {
acc.push_bind(h);
}
FilterValue::RanjIdDesc(r) => {
acc.push_bind(r);
}
FilterValue::Decimal(d) => {
acc.push_bind(d);
}
FilterValue::Interval(i) => {
acc.push_bind(i);
}
#[cfg(feature = "network")]
FilterValue::Inet(addr) => {
acc.push_bind(addr);
}
#[cfg(feature = "network")]
FilterValue::Cidr(cidr) => {
acc.push_bind(cidr);
}
#[cfg(feature = "network")]
FilterValue::Macaddr(mac) => {
acc.push_bind(mac);
}
FilterValue::RangeI32(v) => {
acc.push_bind(v);
}
FilterValue::RangeI64(v) => {
acc.push_bind(v);
}
FilterValue::RangeDecimal(v) => {
acc.push_bind(v);
}
FilterValue::RangeTimestamp(v) => {
acc.push_bind(v);
}
FilterValue::RangeDateTime(v) => {
acc.push_bind(v);
}
FilterValue::RangeDate(v) => {
acc.push_bind(v);
}
FilterValue::Null => {
acc.push_null_literal();
}
FilterValue::ArrayString(v) => {
acc.push_bind(v);
}
FilterValue::ArrayI16(v) => {
acc.push_bind(v);
}
FilterValue::ArrayI32(v) => {
acc.push_bind(v);
}
FilterValue::ArrayI64(v) => {
acc.push_bind(v);
}
FilterValue::ArrayF32(v) => {
acc.push_bind(v);
}
FilterValue::ArrayF64(v) => {
acc.push_bind(v);
}
FilterValue::ArrayBool(v) => {
acc.push_bind(v);
}
FilterValue::ArrayDateTime(v) => {
acc.push_bind(v);
}
FilterValue::ArrayDate(v) => {
acc.push_bind(v);
}
FilterValue::ArrayUuid(v) => {
acc.push_bind(v);
}
FilterValue::ArrayDecimal(v) => {
acc.push_bind(v);
}
FilterValue::ArrayHeerId(v) => {
acc.push_bind(v);
}
FilterValue::ArrayRanjId(v) => {
acc.push_bind(v);
}
FilterValue::ArrayHeerIdDesc(v) => {
acc.push_bind(v);
}
FilterValue::ArrayRanjIdDesc(v) => {
acc.push_bind(v);
}
FilterValue::List(_) | FilterValue::Pair(_, _) => {
unreachable!("push_filter_value called with List/Pair — use emit_leaf")
}
}
}
pub(crate) fn push_filter_value_ref(acc: &mut SqlAccumulator, v: &FilterValue) {
match v {
FilterValue::String(s) => {
acc.push_bind(s.clone());
}
FilterValue::I16(n) => {
acc.push_bind(*n);
}
FilterValue::I32(n) => {
acc.push_bind(*n);
}
FilterValue::I64(n) => {
acc.push_bind(*n);
}
FilterValue::F32(n) => {
acc.push_bind(*n);
}
FilterValue::F64(n) => {
acc.push_bind(*n);
}
FilterValue::Bool(b) => {
acc.push_bind(*b);
}
FilterValue::Timestamp(d) => {
acc.push_bind(*d);
}
FilterValue::DateTime(d) => {
acc.push_bind(*d);
}
FilterValue::Date(d) => {
acc.push_bind(*d);
}
FilterValue::Uuid(u) => {
acc.push_bind(*u);
}
FilterValue::HeerId(h) => {
acc.push_bind(*h);
}
FilterValue::RanjId(r) => {
acc.push_bind(*r);
}
FilterValue::HeerIdDesc(h) => {
acc.push_bind(*h);
}
FilterValue::RanjIdDesc(r) => {
acc.push_bind(*r);
}
FilterValue::Decimal(d) => {
acc.push_bind(*d);
}
FilterValue::Interval(i) => {
acc.push_bind(*i);
}
#[cfg(feature = "network")]
FilterValue::Inet(addr) => {
acc.push_bind(*addr);
}
#[cfg(feature = "network")]
FilterValue::Cidr(cidr) => {
acc.push_bind(*cidr);
}
#[cfg(feature = "network")]
FilterValue::Macaddr(mac) => {
acc.push_bind(*mac);
}
FilterValue::RangeI32(v) => {
acc.push_bind(*v);
}
FilterValue::RangeI64(v) => {
acc.push_bind(*v);
}
FilterValue::RangeDecimal(v) => {
acc.push_bind(*v);
}
FilterValue::RangeTimestamp(v) => {
acc.push_bind(*v);
}
FilterValue::RangeDateTime(v) => {
acc.push_bind(*v);
}
FilterValue::RangeDate(v) => {
acc.push_bind(*v);
}
FilterValue::Null => {
acc.push_null_literal();
}
FilterValue::ArrayString(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayI16(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayI32(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayI64(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayF32(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayF64(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayBool(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayDateTime(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayDate(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayUuid(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayDecimal(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayHeerId(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayRanjId(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayHeerIdDesc(v) => {
acc.push_bind(v.clone());
}
FilterValue::ArrayRanjIdDesc(v) => {
acc.push_bind(v.clone());
}
FilterValue::List(_) | FilterValue::Pair(_, _) => {
unreachable!("push_filter_value_ref called with List/Pair — use emit_leaf_ref")
}
}
}
#[allow(dead_code)]
fn push_list_element(acc: &mut SqlAccumulator, v: FilterValue) {
match v {
FilterValue::Null | FilterValue::List(_) | FilterValue::Pair(_, _) => {
unreachable!("nested/null FilterValue in IN list — typed FieldRef API prevents this")
}
scalar => push_filter_value(acc, scalar),
}
}
fn push_qualified_col(
acc: &mut SqlAccumulator,
col: &'static str,
parent_table: Option<&'static str>,
) {
if let Some(table) = parent_table {
acc.push_sql(table);
acc.push_sql(".");
}
acc.push_sql(col);
}
#[allow(dead_code)]
fn emit_leaf(acc: &mut SqlAccumulator, leaf: Leaf, parent_table: Option<&'static str>) {
let col = leaf.column;
if let Some(tok) = leaf.op.binary_op_token() {
push_qualified_col(acc, col, parent_table);
acc.push_sql(tok);
push_filter_value(acc, leaf.value);
return;
}
match leaf.op {
LookupOp::IsNull => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" IS NULL");
}
LookupOp::IsNotNull => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" IS NOT NULL");
}
LookupOp::IContains => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" ILIKE ");
let s = match leaf.value {
FilterValue::String(s) => s,
_ => unreachable!("IContains requires FilterValue::String"),
};
acc.push_bind(format!("%{}%", escape_like(&s)));
}
LookupOp::IStartsWith => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" ILIKE ");
let s = match leaf.value {
FilterValue::String(s) => s,
_ => unreachable!("IStartsWith requires FilterValue::String"),
};
acc.push_bind(format!("{}%", escape_like(&s)));
}
LookupOp::IEndsWith => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" ILIKE ");
let s = match leaf.value {
FilterValue::String(s) => s,
_ => unreachable!("IEndsWith requires FilterValue::String"),
};
acc.push_bind(format!("%{}", escape_like(&s)));
}
LookupOp::IExact => {
acc.push_sql("LOWER(");
push_qualified_col(acc, col, parent_table);
acc.push_sql(") = LOWER(");
push_filter_value(acc, leaf.value);
acc.push_sql(")");
}
LookupOp::Between => {
let (a, b) = match leaf.value {
FilterValue::Pair(a, b) => (*a, *b),
_ => unreachable!("Between requires FilterValue::Pair"),
};
push_qualified_col(acc, col, parent_table);
acc.push_sql(" BETWEEN ");
push_filter_value(acc, a);
acc.push_sql(" AND ");
push_filter_value(acc, b);
}
LookupOp::In | LookupOp::NotIn => {
let list = match leaf.value {
FilterValue::List(v) => v,
_ => unreachable!("In/NotIn requires FilterValue::List"),
};
if list.is_empty() {
if matches!(leaf.op, LookupOp::In) {
acc.push_sql("FALSE");
} else {
acc.push_sql("TRUE");
}
return;
}
push_qualified_col(acc, col, parent_table);
acc.push_sql(if matches!(leaf.op, LookupOp::In) {
" IN ("
} else {
" NOT IN ("
});
for (i, v) in list.into_iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
push_list_element(acc, v);
}
acc.push_sql(")");
}
LookupOp::Eq
| LookupOp::Neq
| LookupOp::Gt
| LookupOp::Gte
| LookupOp::Lt
| LookupOp::Lte
| LookupOp::Regex
| LookupOp::IRegex => unreachable!("binary-op LookupOp routed past early return"),
}
}
fn emit_leaf_ref(acc: &mut SqlAccumulator, leaf: &Leaf, parent_table: Option<&'static str>) {
let col = leaf.column;
if let Some(tok) = leaf.op.binary_op_token() {
push_qualified_col(acc, col, parent_table);
acc.push_sql(tok);
push_filter_value_ref(acc, &leaf.value);
return;
}
match leaf.op {
LookupOp::IsNull => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" IS NULL");
}
LookupOp::IsNotNull => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" IS NOT NULL");
}
LookupOp::IContains => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" ILIKE ");
let s = match &leaf.value {
FilterValue::String(s) => s,
_ => unreachable!("IContains requires FilterValue::String"),
};
acc.push_bind(format!("%{}%", escape_like(s)));
}
LookupOp::IStartsWith => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" ILIKE ");
let s = match &leaf.value {
FilterValue::String(s) => s,
_ => unreachable!("IStartsWith requires FilterValue::String"),
};
acc.push_bind(format!("{}%", escape_like(s)));
}
LookupOp::IEndsWith => {
push_qualified_col(acc, col, parent_table);
acc.push_sql(" ILIKE ");
let s = match &leaf.value {
FilterValue::String(s) => s,
_ => unreachable!("IEndsWith requires FilterValue::String"),
};
acc.push_bind(format!("%{}", escape_like(s)));
}
LookupOp::IExact => {
acc.push_sql("LOWER(");
push_qualified_col(acc, col, parent_table);
acc.push_sql(") = LOWER(");
push_filter_value_ref(acc, &leaf.value);
acc.push_sql(")");
}
LookupOp::Between => {
let (a, b) = match &leaf.value {
FilterValue::Pair(a, b) => (a.as_ref(), b.as_ref()),
_ => unreachable!("Between requires FilterValue::Pair"),
};
push_qualified_col(acc, col, parent_table);
acc.push_sql(" BETWEEN ");
push_filter_value_ref(acc, a);
acc.push_sql(" AND ");
push_filter_value_ref(acc, b);
}
LookupOp::In | LookupOp::NotIn => {
let list = match &leaf.value {
FilterValue::List(v) => v,
_ => unreachable!("In/NotIn requires FilterValue::List"),
};
if list.is_empty() {
if matches!(leaf.op, LookupOp::In) {
acc.push_sql("FALSE");
} else {
acc.push_sql("TRUE");
}
return;
}
push_qualified_col(acc, col, parent_table);
acc.push_sql(if matches!(leaf.op, LookupOp::In) {
" IN ("
} else {
" NOT IN ("
});
for (i, v) in list.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
push_list_element_ref(acc, v);
}
acc.push_sql(")");
}
LookupOp::Eq
| LookupOp::Neq
| LookupOp::Gt
| LookupOp::Gte
| LookupOp::Lt
| LookupOp::Lte
| LookupOp::Regex
| LookupOp::IRegex => unreachable!("binary-op LookupOp routed past early return"),
}
}
fn push_list_element_ref(acc: &mut SqlAccumulator, v: &FilterValue) {
match v {
FilterValue::Null | FilterValue::List(_) | FilterValue::Pair(_, _) => {
unreachable!("nested/null FilterValue in IN list — typed FieldRef API prevents this")
}
scalar => push_filter_value_ref(acc, scalar),
}
}
#[cfg(test)]
mod phase85_array_in_regression_tests {
use super::*;
#[test]
fn array_values_in_in_list_bind_instead_of_panicking() {
let mut acc = SqlAccumulator::new("");
let value = FilterValue::ArrayI32(vec![1, 2, 3]);
push_list_element_ref(&mut acc, &value);
assert_eq!(acc.sql(), "$1");
assert_eq!(acc.bind_count(), 1);
}
}
pub(crate) fn emit_condition(
acc: &mut SqlAccumulator,
c: &Condition,
parent_table: Option<&'static str>,
) -> Result<(), PortablePredicateError> {
match c {
Condition::True => {
acc.push_sql("TRUE");
Ok(())
}
Condition::Leaf(l) => {
emit_leaf_ref(acc, l, parent_table);
Ok(())
}
Condition::Not(inner) => {
acc.push_sql("NOT (");
emit_condition(acc, inner, parent_table)?;
acc.push_sql(")");
Ok(())
}
Condition::And(parts) => {
if parts.is_empty() {
acc.push_sql("TRUE");
return Ok(());
}
acc.push_sql("(");
for (i, p) in parts.iter().enumerate() {
if i > 0 {
acc.push_sql(" AND ");
}
emit_condition(acc, p, parent_table)?;
}
acc.push_sql(")");
Ok(())
}
Condition::Or(parts) => {
if parts.is_empty() {
acc.push_sql("FALSE");
return Ok(());
}
acc.push_sql("(");
for (i, p) in parts.iter().enumerate() {
if i > 0 {
acc.push_sql(" OR ");
}
emit_condition(acc, p, parent_table)?;
}
acc.push_sql(")");
Ok(())
}
Condition::Expr(expr) => {
let ctx = match parent_table {
Some(t) => SqlEmitContext::joined(t),
None => SqlEmitContext::root(),
};
crate::expr::sql::emit_expr(acc, &expr.node, ctx)
}
Condition::ArrayContains(leaf) => {
push_qualified_col(acc, leaf.column, parent_table);
acc.push_sql(" @> ");
push_filter_value_ref(acc, &leaf.values);
Ok(())
}
Condition::ArrayContainedBy(leaf) => {
push_qualified_col(acc, leaf.column, parent_table);
acc.push_sql(" <@ ");
push_filter_value_ref(acc, &leaf.values);
Ok(())
}
Condition::ArrayOverlap(leaf) => {
push_qualified_col(acc, leaf.column, parent_table);
acc.push_sql(" && ");
push_filter_value_ref(acc, &leaf.values);
Ok(())
}
Condition::RangePredicate(leaf) => {
push_qualified_col(acc, leaf.column(), parent_table);
acc.push_sql(leaf.op().sql_token());
push_filter_value_ref(acc, leaf.value());
if let Some(cast) = leaf.rhs_element_cast() {
acc.push_sql("::");
acc.push_sql(cast);
}
Ok(())
}
Condition::JsonbPath(leaf) => {
emit_jsonb_path_leaf_ref(acc, leaf, parent_table);
Ok(())
}
Condition::RawSql(s) => {
acc.push_sql("(");
acc.push_sql(s.as_str());
acc.push_sql(")");
Ok(())
}
}
}
fn emit_jsonb_path_leaf_ref(
acc: &mut SqlAccumulator,
leaf: &crate::jsonb::path::JsonbPathLeaf,
parent_table: Option<&'static str>,
) {
fn build_lhs(
acc: &mut SqlAccumulator,
column: &'static str,
path: &'static str,
cast: Option<&'static str>,
parent_table: Option<&'static str>,
) {
let segments: Vec<&str> = path.split('.').collect();
acc.push_sql("(");
if let Some(table) = parent_table {
acc.push_sql(table);
acc.push_sql(".");
}
acc.push_sql(column);
for (i, seg) in segments.iter().enumerate() {
if i == segments.len() - 1 {
acc.push_sql("->>'");
acc.push_sql(seg);
acc.push_sql("'");
} else {
acc.push_sql("->'");
acc.push_sql(seg);
acc.push_sql("'");
}
}
acc.push_sql(")");
if let Some(c) = cast {
acc.push_sql(c);
}
}
if matches!(leaf.op, LookupOp::Regex | LookupOp::IRegex) {
unreachable!(
"Regex / IRegex not supported on JsonbPathLeaf: {:?}",
leaf.op
);
}
if let Some(tok) = leaf.op.binary_op_token() {
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(tok);
push_filter_value_ref(acc, &leaf.value);
return;
}
match leaf.op {
LookupOp::IsNull => {
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(" IS NULL");
}
LookupOp::IsNotNull => {
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(" IS NOT NULL");
}
LookupOp::In => {
let list = match &leaf.value {
FilterValue::List(v) => v,
_ => unreachable!("JsonbPath In requires FilterValue::List"),
};
if list.is_empty() {
acc.push_sql("FALSE");
return;
}
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(" IN (");
for (i, v) in list.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
push_list_element_ref(acc, v);
}
acc.push_sql(")");
}
_ => {
unreachable!("unsupported LookupOp in JsonbPathLeaf: {:?}", leaf.op)
}
}
}
#[allow(dead_code)]
fn emit_jsonb_path_leaf(
acc: &mut SqlAccumulator,
leaf: crate::jsonb::path::JsonbPathLeaf,
parent_table: Option<&'static str>,
) {
use LookupOp::{In, IsNotNull, IsNull};
fn build_lhs(
acc: &mut SqlAccumulator,
column: &'static str,
path: &'static str,
cast: Option<&'static str>,
parent_table: Option<&'static str>,
) {
let segments: Vec<&str> = path.split('.').collect();
acc.push_sql("(");
if let Some(table) = parent_table {
acc.push_sql(table);
acc.push_sql(".");
}
acc.push_sql(column);
for (i, seg) in segments.iter().enumerate() {
if i == segments.len() - 1 {
acc.push_sql("->>'");
acc.push_sql(seg);
acc.push_sql("'");
} else {
acc.push_sql("->'");
acc.push_sql(seg);
acc.push_sql("'");
}
}
acc.push_sql(")");
if let Some(c) = cast {
acc.push_sql(c);
}
}
if matches!(leaf.op, LookupOp::Regex | LookupOp::IRegex) {
unreachable!(
"Regex / IRegex not supported on JsonbPathLeaf: {:?}",
leaf.op
);
}
if let Some(tok) = leaf.op.binary_op_token() {
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(tok);
push_filter_value(acc, leaf.value);
return;
}
match leaf.op {
IsNull => {
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(" IS NULL");
}
IsNotNull => {
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(" IS NOT NULL");
}
In => {
let list = match leaf.value {
FilterValue::List(v) => v,
_ => unreachable!("JsonbPath In requires FilterValue::List"),
};
if list.is_empty() {
acc.push_sql("FALSE");
return;
}
build_lhs(acc, leaf.column, leaf.path, leaf.cast, parent_table);
acc.push_sql(" IN (");
for (i, v) in list.into_iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
push_list_element(acc, v);
}
acc.push_sql(")");
}
_ => {
unreachable!("unsupported LookupOp in JsonbPathLeaf: {:?}", leaf.op)
}
}
}
pub(crate) fn emit_q<T: Model>(
acc: &mut SqlAccumulator,
q: &Q<T>,
ctx: SqlEmitContext,
) -> Result<(), PortablePredicateError> {
let parent_table = ctx.parent_table();
match q {
Q::Portable(predicate) => {
crate::query::portable::emit_portable_predicate::<T>(acc, predicate, ctx)
}
Q::Ilike(field, pattern) => {
push_qualified_col(acc, field.column(), parent_table);
acc.push_sql(" ILIKE ");
acc.push_bind(format!("%{}%", escape_like(pattern)));
Ok(())
}
Q::JsonbPath(leaf) => {
emit_jsonb_path_leaf_ref(acc, leaf, parent_table);
Ok(())
}
Q::Regex(field, pattern, true) => {
push_qualified_col(acc, field.column(), parent_table);
acc.push_sql(" ~ ");
acc.push_bind(pattern.clone());
Ok(())
}
Q::Regex(field, pattern, false) => {
push_qualified_col(acc, field.column(), parent_table);
acc.push_sql(" ~* ");
acc.push_bind(pattern.clone());
Ok(())
}
Q::Expression(expr) => crate::expr::sql::emit_expr(acc, &expr.node, ctx),
Q::Array(ArrayPredicate::Contains(leaf, _)) => {
push_qualified_col(acc, leaf.column, parent_table);
acc.push_sql(" @> ");
push_filter_value_ref(acc, &leaf.values);
Ok(())
}
Q::Array(ArrayPredicate::ContainedBy(leaf, _)) => {
push_qualified_col(acc, leaf.column, parent_table);
acc.push_sql(" <@ ");
push_filter_value_ref(acc, &leaf.values);
Ok(())
}
Q::Array(ArrayPredicate::Overlap(leaf, _)) => {
push_qualified_col(acc, leaf.column, parent_table);
acc.push_sql(" && ");
push_filter_value_ref(acc, &leaf.values);
Ok(())
}
Q::Condition(c) => emit_condition(acc, c, parent_table),
Q::Compound { op, parts } => {
if parts.is_empty() {
acc.push_sql(match op {
CompoundOp::And => "TRUE",
CompoundOp::Or => "FALSE",
});
return Ok(());
}
acc.push_sql("(");
let sep = match op {
CompoundOp::And => " AND ",
CompoundOp::Or => " OR ",
};
for (i, p) in parts.iter().enumerate() {
if i > 0 {
acc.push_sql(sep);
}
emit_q::<T>(acc, p, ctx)?;
}
acc.push_sql(")");
Ok(())
}
Q::Xor(a, b) => {
acc.push_sql("(((NOT (");
emit_q::<T>(acc, a, ctx)?;
acc.push_sql(")) AND (");
emit_q::<T>(acc, b, ctx)?;
acc.push_sql(")) OR ((");
emit_q::<T>(acc, a, ctx)?;
acc.push_sql(") AND (NOT (");
emit_q::<T>(acc, b, ctx)?;
acc.push_sql("))))");
Ok(())
}
Q::Negated(inner) => {
acc.push_sql("NOT (");
emit_q::<T>(acc, inner, ctx)?;
acc.push_sql(")");
Ok(())
}
#[allow(unreachable_patterns)]
_ => Err(PortablePredicateError::CacheInvalidNode {
kind: "Q::<unknown>",
}),
}
}
pub(crate) fn q_is_vacuously_true<T: Model>(q: &Q<T>) -> bool {
match q {
Q::Portable(p) => match p.inner_ref() {
sassi::BasicPredicate::True => true,
sassi::BasicPredicate::And(parts) => parts
.iter()
.all(|c| matches!(c, sassi::BasicPredicate::True)),
_ => false,
},
Q::Condition(c) => c.is_vacuously_true(),
Q::Compound {
op: CompoundOp::And,
parts,
} => parts.iter().all(q_is_vacuously_true),
Q::Negated(inner) => match inner.as_ref() {
Q::Portable(p) => matches!(p.inner_ref(), sassi::BasicPredicate::False),
Q::Compound {
op: CompoundOp::Or,
parts,
} => parts.is_empty(),
_ => false,
},
_ => false,
}
}
fn push_where<T: Model>(
acc: &mut SqlAccumulator,
qs: &QuerySet<T>,
) -> Result<(), PortablePredicateError> {
push_where_qualified(acc, qs, None)
}
pub(crate) fn push_where_with_ctx<T: Model>(
acc: &mut SqlAccumulator,
qs: &QuerySet<T>,
ctx: SqlEmitContext,
) -> Result<(), PortablePredicateError> {
if q_is_vacuously_true(&qs.condition) {
return Ok(());
}
acc.push_sql(" WHERE ");
emit_q::<T>(acc, &qs.condition, ctx)
}
fn push_where_qualified<T: Model>(
acc: &mut SqlAccumulator,
qs: &QuerySet<T>,
parent_table: Option<&'static str>,
) -> Result<(), PortablePredicateError> {
let ctx = match parent_table {
Some(t) => SqlEmitContext::joined(t),
None => SqlEmitContext::root(),
};
push_where_with_ctx(acc, qs, ctx)
}
fn push_tail<T: Model>(
acc: &mut SqlAccumulator,
qs: &QuerySet<T>,
) -> Result<(), PortablePredicateError> {
push_tail_qualified(acc, qs, None)
}
pub(crate) fn push_tail_with_ctx<T: Model>(
acc: &mut SqlAccumulator,
qs: &QuerySet<T>,
ctx: SqlEmitContext,
) -> Result<(), PortablePredicateError> {
push_where_with_ctx(acc, qs, ctx)?;
if !qs.ordering.is_empty() {
acc.push_sql(" ORDER BY ");
for (i, o) in qs.ordering.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
o.emit(acc, ctx.parent_table());
}
}
if let Some(n) = qs.limit {
acc.push_sql(" LIMIT ");
acc.push_bind(n);
}
if let Some(n) = qs.offset {
acc.push_sql(" OFFSET ");
acc.push_bind(n);
}
qs.lock.push_tail(acc);
Ok(())
}
pub(crate) fn push_tail_qualified<T: Model>(
acc: &mut SqlAccumulator,
qs: &QuerySet<T>,
parent_table: Option<&'static str>,
) -> Result<(), PortablePredicateError> {
let ctx = match parent_table {
Some(t) => SqlEmitContext::joined(t),
None => SqlEmitContext::root(),
};
push_tail_with_ctx(acc, qs, ctx)
}
fn push_grouped_tail(
acc: &mut SqlAccumulator,
having: Option<&crate::expr::node::ExprNode>,
order: &[crate::query::order::OrderExpr],
limit: Option<u64>,
offset: Option<u64>,
) -> Result<(), PortablePredicateError> {
if let Some(h) = having {
acc.push_sql(" HAVING ");
crate::expr::sql::emit_expr(acc, h, SqlEmitContext::root())?;
}
if !order.is_empty() {
acc.push_sql(" ORDER BY ");
for (i, o) in order.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
o.emit(acc, None);
}
}
if let Some(n) = limit {
acc.push_sql(" LIMIT ");
acc.push_bind(n as i64);
}
if let Some(n) = offset {
acc.push_sql(" OFFSET ");
acc.push_bind(n as i64);
}
Ok(())
}
pub(crate) fn build_select<T: Model + FromPgRow>(
qs: &QuerySet<T>,
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("");
match &qs.distinct {
DistinctMode::None => {
acc.push_sql("SELECT ");
acc.push_sql(<T as FromPgRow>::COLUMN_LIST);
acc.push_sql(" FROM ");
}
DistinctMode::Plain => {
acc.push_sql("SELECT DISTINCT ");
acc.push_sql(<T as FromPgRow>::COLUMN_LIST);
acc.push_sql(" FROM ");
}
DistinctMode::On(cols) => {
acc.push_sql("SELECT DISTINCT ON (");
acc.push_csv(cols.iter().copied());
acc.push_sql(") ");
acc.push_sql(<T as FromPgRow>::COLUMN_LIST);
acc.push_sql(" FROM ");
}
}
acc.push_sql(T::table_name());
push_tail(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_select_joined<T: Model>(
qs: &QuerySet<T>,
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("");
let col_list = crate::relation::select_related::select_columns::<T>(&qs.select_related_paths);
let parent_table: Option<&'static str> = Some(T::table_name());
match &qs.distinct {
DistinctMode::None => {
acc.push_sql("SELECT ");
acc.push_sql(&col_list);
acc.push_sql(" FROM ");
}
DistinctMode::Plain => {
acc.push_sql("SELECT DISTINCT ");
acc.push_sql(&col_list);
acc.push_sql(" FROM ");
}
DistinctMode::On(cols) => {
acc.push_sql("SELECT DISTINCT ON (");
for (i, c) in cols.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql(T::table_name());
acc.push_sql(".");
acc.push_sql(c);
}
acc.push_sql(") ");
acc.push_sql(&col_list);
acc.push_sql(" FROM ");
}
}
acc.push_sql(T::table_name());
crate::relation::select_related::push_joins::<T>(&mut acc, &qs.select_related_paths);
push_tail_qualified(&mut acc, qs, parent_table)?;
Ok(acc)
}
fn emit_aggregate_inner(
acc: &mut SqlAccumulator,
agg: &crate::expr::node::ExprNode,
default_window: Option<&'static str>,
) -> Result<(), PortablePredicateError> {
let shape = spatial_emission_shape(agg);
let (cast_to, window) = match agg {
crate::expr::node::ExprNode::Aggregate {
cast_to, window, ..
} => (*cast_to, window.as_ref()),
_ => (None, None),
};
let has_window_clause = window.is_some() || default_window.is_some();
let emit_window_clause = |acc: &mut SqlAccumulator| match window {
Some(ws) => ws.emit(acc),
None => {
if let Some(s) = default_window {
acc.push_sql(s);
}
}
};
match (shape, has_window_clause, cast_to) {
(
Some(SpatialShape {
suffix,
wrapped: true,
}),
true,
_,
) => {
crate::expr::sql::emit_expr(acc, agg, SqlEmitContext::root())?;
let popped_cast = acc.pop_sql_suffix(suffix);
debug_assert!(
popped_cast,
"wrapped spatial bare emission must end with {suffix}"
);
let popped_close = acc.pop_sql_suffix(")");
debug_assert!(
popped_close,
"wrapped spatial bare emission must end with `)` after popping cast"
);
emit_window_clause(acc);
acc.push_sql(")");
acc.push_sql(suffix);
}
(
Some(SpatialShape {
suffix,
wrapped: false,
}),
true,
_,
) => {
let has_filter = matches!(
agg,
crate::expr::node::ExprNode::Aggregate {
filter: Some(_),
..
}
);
if has_filter {
crate::expr::sql::emit_expr(acc, agg, SqlEmitContext::root())?;
let popped_cast = acc.pop_sql_suffix(suffix);
debug_assert!(
popped_cast,
"unwrapped spatial bare emission must end with {suffix}"
);
let popped_close = acc.pop_sql_suffix(")");
debug_assert!(
popped_close,
"unwrapped spatial-with-filter bare emission must end with `)` \
after popping cast (FILTER parens from emit_spatial_unary_agg)"
);
emit_window_clause(acc);
acc.push_sql(")");
acc.push_sql(suffix);
} else {
acc.push_sql("(");
crate::expr::sql::emit_expr(acc, agg, SqlEmitContext::root())?;
let popped_cast = acc.pop_sql_suffix(suffix);
debug_assert!(
popped_cast,
"unwrapped spatial bare emission must end with {suffix}"
);
emit_window_clause(acc);
acc.push_sql(")");
acc.push_sql(suffix);
}
}
(Some(_), false, _) => {
crate::expr::sql::emit_expr(acc, agg, SqlEmitContext::root())?;
}
(None, _, Some(ty)) => {
acc.push_sql("(");
crate::expr::sql::emit_expr(acc, agg, SqlEmitContext::root())?;
emit_window_clause(acc);
acc.push_sql(")::");
acc.push_sql(ty);
}
(None, _, None) => {
crate::expr::sql::emit_expr(acc, agg, SqlEmitContext::root())?;
emit_window_clause(acc);
}
}
Ok(())
}
struct SpatialShape {
suffix: &'static str,
wrapped: bool,
}
fn spatial_emission_shape(agg: &crate::expr::node::ExprNode) -> Option<SpatialShape> {
use crate::expr::node::ExprNode;
match agg {
ExprNode::Aggregate { op, .. } => {
let suffix = crate::expr::sql::outer_cast_suffix(op)?;
#[cfg(feature = "spatial")]
let wrapped = matches!(
op,
crate::expr::node::AggOp::SpatialCentroid
| crate::expr::node::AggOp::SpatialConvexHull
);
#[cfg(not(feature = "spatial"))]
let wrapped = false;
Some(SpatialShape { suffix, wrapped })
}
_ => None,
}
}
pub(crate) fn emit_aggregate_with_cast(
acc: &mut SqlAccumulator,
agg: &crate::expr::node::ExprNode,
) -> Result<(), PortablePredicateError> {
emit_aggregate_inner(acc, agg, None)
}
pub(crate) fn emit_aggregate_with_window_and_cast(
acc: &mut SqlAccumulator,
agg: &crate::expr::node::ExprNode,
) -> Result<(), PortablePredicateError> {
emit_aggregate_inner(acc, agg, Some(" OVER ()"))
}
pub(crate) fn build_aggregate_select<T: Model>(
qs: &QuerySet<T>,
agg: &crate::expr::node::ExprNode,
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("SELECT ");
emit_aggregate_with_cast(&mut acc, agg)?;
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
push_where(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_select_with_annotations<T, F>(
qs: &QuerySet<T>,
push_columns: F,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model + FromPgRow,
F: FnOnce(&mut SqlAccumulator),
{
let mut acc = SqlAccumulator::new("");
acc.push_sql("SELECT ");
for (i, col) in <T as FromPgRow>::COLUMNS.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql("t.");
acc.push_sql(col);
}
push_columns(&mut acc);
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" AS t");
push_tail(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_annotated_select_for_fetch<T, F>(
qs: &QuerySet<T>,
push_columns: F,
qualify: Option<&crate::expr::QualifyCondition>,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model + FromPgRow,
F: FnOnce(&mut SqlAccumulator),
{
let inner = build_select_with_annotations(qs, push_columns)?;
let Some(qualify) = qualify else {
return Ok(inner);
};
let mut wrapped = SqlAccumulator::new("SELECT * FROM (");
wrapped.extend_with(inner);
wrapped.push_sql(") AS __djogi_q WHERE ");
qualify.push_outer_where(&mut wrapped);
Ok(wrapped)
}
#[cfg(feature = "spatial")]
pub(crate) fn build_spatial_row_select_with_annotations_for_fetch<T, F>(
qs: &QuerySet<T>,
push_columns: F,
qualify: Option<&crate::expr::QualifyCondition>,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model + FromPgRow,
F: FnOnce(&mut SqlAccumulator),
{
let inner = build_spatial_row_select_with_annotations(qs, push_columns)?;
let Some(qualify) = qualify else {
return Ok(inner);
};
let mut wrapped = SqlAccumulator::new("SELECT * FROM (");
wrapped.extend_with(inner);
wrapped.push_sql(") AS __djogi_q WHERE ");
qualify.push_outer_where(&mut wrapped);
Ok(wrapped)
}
#[cfg(feature = "spatial")]
fn build_spatial_row_select_with_annotations<T, F>(
qs: &QuerySet<T>,
push_columns: F,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model + FromPgRow,
F: FnOnce(&mut SqlAccumulator),
{
let mut acc = SqlAccumulator::new("SELECT ");
let desc = T::descriptor();
for (i, col) in <T as FromPgRow>::COLUMNS.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql("t.");
acc.push_sql(col);
if desc.fields.iter().any(|f| {
f.name == *col
&& matches!(
f.sql_type,
crate::descriptor::FieldSqlType::Geography { .. }
)
}) {
acc.push_sql("::geometry AS ");
acc.push_sql(col);
}
}
push_columns(&mut acc);
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" AS t");
push_tail(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_grouped_annotated_select<T, K, A>(
gaq: &crate::query::grouped::GroupedAnnotatedQuerySet<T, K, A>,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model,
K: crate::query::grouped::IntoGroupKeyTuple,
A: crate::query::annotate::IntoAggregateTuple,
{
#[cfg(feature = "spatial")]
if let Some(ref src) = gaq.spatial_source {
return match src {
crate::query::grouped::SpatialGroupSource::Join(spec) => {
build_spatial_join_grouped_select(gaq, spec)
}
crate::query::grouped::SpatialGroupSource::Cluster(spec) => {
build_cluster_grouped_select(gaq, spec)
}
crate::query::grouped::SpatialGroupSource::Geohash(spec) => {
build_geohash_grouped_select(gaq, spec)
}
};
}
let mut acc = SqlAccumulator::new("SELECT ");
gaq.keys.push_select_columns(&mut acc);
let has_key_columns = acc.sql() != "SELECT ";
gaq.aggregates
.push_columns_bare_after(&mut acc, has_key_columns);
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" AS t");
push_where(&mut acc, &gaq.qs)?;
acc.push_sql(" GROUP BY ");
match gaq.grouping {
crate::query::grouped::GroupingMode::Plain => {
gaq.keys.push_group_by_columns(&mut acc);
}
crate::query::grouped::GroupingMode::Rollup => {
acc.push_sql("ROLLUP (");
gaq.keys.push_group_by_columns(&mut acc);
acc.push_sql(")");
}
crate::query::grouped::GroupingMode::Cube => {
acc.push_sql("CUBE (");
gaq.keys.push_group_by_columns(&mut acc);
acc.push_sql(")");
}
crate::query::grouped::GroupingMode::Sets(ref sets) => {
acc.push_sql("GROUPING SETS (");
for (i, set) in sets.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql("(");
acc.push_csv(set.iter().copied());
acc.push_sql(")");
}
acc.push_sql(")");
}
}
push_grouped_tail(
&mut acc,
gaq.having.as_ref(),
&gaq.order,
gaq.limit,
gaq.offset,
)?;
Ok(acc)
}
#[cfg(feature = "spatial")]
pub(crate) fn build_spatial_join_grouped_select<T, K, A>(
gaq: &crate::query::grouped::GroupedAnnotatedQuerySet<T, K, A>,
spec: &crate::query::spatial_grouping::SpatialJoinSpec,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model,
K: crate::query::grouped::IntoGroupKeyTuple,
A: crate::query::annotate::IntoAggregateTuple,
{
let mut acc = SqlAccumulator::new("SELECT ");
gaq.keys.push_select_columns(&mut acc);
gaq.aggregates.push_columns_bare(&mut acc);
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" AS t");
acc.push_sql(" LEFT JOIN ");
acc.push_sql(spec.r_table);
acc.push_sql(" AS r ON ST_Covers(r.");
acc.push_sql(spec.r_geo_col);
acc.push_sql(", t.");
acc.push_sql(spec.t_geo_col);
acc.push_sql(")");
push_where_qualified(&mut acc, &gaq.qs, Some("t"))?;
acc.push_sql(" GROUP BY ");
match gaq.grouping {
crate::query::grouped::GroupingMode::Plain => {
gaq.keys.push_group_by_columns(&mut acc);
}
_ => {
gaq.keys.push_group_by_columns(&mut acc);
}
}
push_grouped_tail(
&mut acc,
gaq.having.as_ref(),
&gaq.order,
gaq.limit,
gaq.offset,
)?;
Ok(acc)
}
#[cfg(feature = "spatial")]
pub(crate) fn build_cluster_grouped_select<T, K, A>(
gaq: &crate::query::grouped::GroupedAnnotatedQuerySet<T, K, A>,
spec: &crate::query::spatial_grouping::ClusterSpec,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model,
K: crate::query::grouped::IntoGroupKeyTuple,
A: crate::query::annotate::IntoAggregateTuple,
{
let mut acc = SqlAccumulator::new("SELECT cluster_id");
gaq.aggregates.push_columns_bare(&mut acc);
acc.push_sql(" FROM (SELECT t.*, ST_ClusterDBSCAN(t.");
acc.push_sql(spec.t_geo_col);
acc.push_sql("::geometry, ");
acc.push_bind(spec.eps_degrees);
acc.push_sql(", ");
acc.push_bind(spec.minpoints);
acc.push_sql(") OVER () AS cluster_id FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" AS t");
push_where(&mut acc, &gaq.qs)?;
acc.push_sql(") AS t");
acc.push_sql(" GROUP BY cluster_id");
push_grouped_tail(
&mut acc,
gaq.having.as_ref(),
&gaq.order,
gaq.limit,
gaq.offset,
)?;
Ok(acc)
}
#[cfg(feature = "spatial")]
pub(crate) fn build_geohash_grouped_select<T, K, A>(
gaq: &crate::query::grouped::GroupedAnnotatedQuerySet<T, K, A>,
spec: &crate::query::spatial_grouping::GeohashSpec,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model,
K: crate::query::grouped::IntoGroupKeyTuple,
A: crate::query::annotate::IntoAggregateTuple,
{
let mut acc = SqlAccumulator::new("SELECT ");
acc.push_sql("ST_GeoHash(t.");
acc.push_sql(spec.t_geo_col);
acc.push_sql("::geometry, ");
acc.push_bind(spec.precision);
acc.push_sql(") AS geohash");
gaq.aggregates.push_columns_bare(&mut acc);
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" AS t");
push_where(&mut acc, &gaq.qs)?;
acc.push_sql(" GROUP BY geohash");
push_grouped_tail(
&mut acc,
gaq.having.as_ref(),
&gaq.order,
gaq.limit,
gaq.offset,
)?;
Ok(acc)
}
pub(crate) fn build_count<T: Model>(
qs: &QuerySet<T>,
) -> Result<SqlAccumulator, PortablePredicateError> {
match &qs.distinct {
DistinctMode::None => {
let mut acc = SqlAccumulator::new("SELECT COUNT(*) FROM ");
acc.push_sql(T::table_name());
push_where(&mut acc, qs)?;
Ok(acc)
}
DistinctMode::Plain => {
let mut acc = SqlAccumulator::new("SELECT COUNT(*) FROM (SELECT DISTINCT * FROM ");
acc.push_sql(T::table_name());
push_where(&mut acc, qs)?;
acc.push_sql(") AS sub");
Ok(acc)
}
DistinctMode::On(cols) => {
let mut acc = SqlAccumulator::new("SELECT COUNT(*) FROM (SELECT DISTINCT ON (");
acc.push_csv(cols.iter().copied());
acc.push_sql(") * FROM ");
acc.push_sql(T::table_name());
push_where(&mut acc, qs)?;
acc.push_sql(" ORDER BY ");
acc.push_csv(cols.iter().copied());
for o in qs.ordering.iter() {
acc.push_sql(", ");
o.emit(&mut acc, None);
}
acc.push_sql(") AS sub");
Ok(acc)
}
}
}
pub(crate) fn build_exists<T: Model>(
qs: &QuerySet<T>,
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("SELECT EXISTS(SELECT 1 FROM ");
acc.push_sql(T::table_name());
push_where(&mut acc, qs)?;
acc.push_sql(" LIMIT 1)");
Ok(acc)
}
pub(crate) fn build_update<T: Model>(
qs: &QuerySet<T>,
assignments: &[crate::query::update::UpdateAssignment],
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("UPDATE ");
acc.push_sql(T::table_name());
acc.push_sql(" SET ");
for (i, a) in assignments.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql(a.column());
acc.push_sql(" = ");
match a.value() {
crate::query::update::AssignmentValue::Literal(v) => {
push_filter_value(&mut acc, v.clone());
}
crate::query::update::AssignmentValue::Expr(node) => {
crate::expr::sql::emit_expr(&mut acc, node, SqlEmitContext::root())?;
}
}
}
acc.push_sql(", updated_at = now()");
push_where(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_delete<T: Model>(
qs: &QuerySet<T>,
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("DELETE FROM ");
acc.push_sql(T::table_name());
push_where(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_merge<S: Model + FromPgRow, T: Model>(
source: &QuerySet<S>,
on: &[MergeOnEq<S, T>],
branches: &[MergeBranch<S, T>],
returning: Option<()>,
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("MERGE INTO ");
acc.push_sql(T::table_name());
acc.push_sql(" AS ");
acc.push_sql(TGT_ALIAS);
acc.push_sql(" ");
let _ = returning;
acc.push_sql("USING (");
let source_acc = build_select(source)?;
acc.extend_with(source_acc);
acc.push_sql(") AS ");
acc.push_sql(SRC_ALIAS);
acc.push_sql(" ");
acc.push_sql("ON ");
for (i, cond) in on.iter().enumerate() {
if i > 0 {
acc.push_sql(" AND ");
}
acc.push_sql(TGT_ALIAS);
acc.push_sql(".");
acc.push_sql(cond.target_col);
acc.push_sql(" = ");
acc.push_sql(SRC_ALIAS);
acc.push_sql(".");
acc.push_sql(cond.source_col);
}
for branch in branches {
acc.push_sql("\nWHEN ");
match branch.match_kind {
MergeMatchKind::Matched => acc.push_sql("MATCHED"),
MergeMatchKind::NotMatchedByTarget => acc.push_sql("NOT MATCHED"),
MergeMatchKind::NotMatchedBySource => acc.push_sql("NOT MATCHED BY SOURCE"),
}
if let Some(cond) = &branch.condition {
acc.push_sql(" AND ");
crate::expr::sql::emit_expr(&mut acc, &cond.node, SqlEmitContext::joined(TGT_ALIAS))?;
}
acc.push_sql(" THEN ");
match &branch.action {
MergeAction::Update(updates) => {
acc.push_sql("UPDATE SET ");
acc.push_sql("updated_at = now()");
for update in updates {
acc.push_sql(", ");
acc.push_sql(update.target_col);
acc.push_sql(" = ");
match &update.value {
MergeValue::Literal(v, _) => push_filter_value(&mut acc, v.clone()),
MergeValue::SourceField(col, _) => {
acc.push_sql(SRC_ALIAS);
acc.push_sql(".");
acc.push_sql(col);
}
MergeValue::TargetExpr(node, _) => {
crate::expr::sql::emit_expr(
&mut acc,
node,
SqlEmitContext::joined(TGT_ALIAS),
)?;
}
}
}
}
MergeAction::Delete => {
acc.push_sql("DELETE");
}
MergeAction::Insert(columns) => {
acc.push_sql("INSERT (");
for (i, col) in columns.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql(col.target_col);
}
acc.push_sql(") VALUES (");
for (i, col) in columns.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
match &col.value {
MergeValue::Literal(v, _) => push_filter_value(&mut acc, v.clone()),
MergeValue::SourceField(scol, _) => {
acc.push_sql(SRC_ALIAS);
acc.push_sql(".");
acc.push_sql(scol);
}
MergeValue::TargetExpr(node, _) => {
crate::expr::sql::emit_expr(
&mut acc,
node,
SqlEmitContext::joined(TGT_ALIAS),
)?;
}
}
}
acc.push_sql(")");
}
MergeAction::_Marker(_) => unreachable!("MergeAction::_Marker is a type marker only"),
}
}
Ok(acc)
}
fn push_old_new_returning_projection<T: FromPgRow>(acc: &mut SqlAccumulator, include_new: bool) {
if include_new {
acc.push_sql(" RETURNING WITH (OLD AS __djogi_old, NEW AS __djogi_new)");
} else {
acc.push_sql(" RETURNING WITH (OLD AS __djogi_old)");
}
let mut first = true;
for (idx, col) in T::COLUMNS.iter().enumerate() {
if first {
acc.push_sql(" ");
first = false;
} else {
acc.push_sql(", ");
}
acc.push_sql("__djogi_old.");
acc.push_sql(col);
let old_alias = joined_alias_for_prefix("__djogi_old__", idx, col);
acc.push_sql(" AS \"");
acc.push_sql(&old_alias);
acc.push_sql("\"");
}
if include_new {
for (idx, col) in T::COLUMNS.iter().enumerate() {
acc.push_sql(", __djogi_new.");
acc.push_sql(col);
let new_alias = joined_alias_for_prefix("__djogi_new__", idx, col);
acc.push_sql(" AS \"");
acc.push_sql(&new_alias);
acc.push_sql("\"");
}
}
}
pub(crate) fn build_update_returning_pairs<T>(
qs: &QuerySet<T>,
assignments: &[crate::query::update::UpdateAssignment],
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model + FromPgRow,
{
let mut acc = build_update(qs, assignments)?;
push_old_new_returning_projection::<T>(&mut acc, true);
Ok(acc)
}
pub(crate) fn build_update_returning_ids<T: Model>(
qs: &QuerySet<T>,
assignments: &[crate::query::update::UpdateAssignment],
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = build_update(qs, assignments)?;
let pk_column = T::descriptor()
.pk_column()
.expect("Model implementations with CRUD support must expose a primary-key column");
acc.push_sql(" RETURNING ");
acc.push_sql(pk_column);
Ok(acc)
}
pub(crate) fn build_delete_returning<T>(
qs: &QuerySet<T>,
) -> Result<SqlAccumulator, PortablePredicateError>
where
T: Model + FromPgRow,
{
let mut acc = build_delete(qs)?;
push_old_new_returning_projection::<T>(&mut acc, false);
Ok(acc)
}
pub(crate) fn build_insert_select<S: Model, T: Model>(
qs: &QuerySet<S>,
columns: &[crate::query::insert_select::InsertSelectColumn<S, T>],
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = SqlAccumulator::new("INSERT INTO ");
acc.push_sql(T::table_name());
acc.push_sql(" (");
for (i, col) in columns.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
acc.push_sql(col.target_column());
}
acc.push_sql(") SELECT ");
for (i, col) in columns.iter().enumerate() {
if i > 0 {
acc.push_sql(", ");
}
crate::expr::sql::emit_expr(&mut acc, col.source(), SqlEmitContext::root())?;
}
acc.push_sql(" FROM ");
acc.push_sql(S::table_name());
push_tail(&mut acc, qs)?;
Ok(acc)
}
pub(crate) fn build_insert_select_returning<S: Model, T: Model + FromPgRow>(
qs: &QuerySet<S>,
columns: &[crate::query::insert_select::InsertSelectColumn<S, T>],
) -> Result<SqlAccumulator, PortablePredicateError> {
let mut acc = build_insert_select::<S, T>(qs, columns)?;
acc.push_sql(" RETURNING ");
acc.push_sql(<T as FromPgRow>::COLUMN_LIST);
Ok(acc)
}
pub(crate) fn assert_no_alias_collision(sql: &str) -> Result<(), crate::DjogiError> {
let after_select = if let Some(s) = sql.strip_prefix("SELECT ") {
s
} else if let Some(i) = sql.find("SELECT ") {
&sql[i + "SELECT ".len()..]
} else {
return Ok(()); };
let from_idx = match after_select.find(" FROM ") {
Some(i) => i,
None => return Ok(()),
};
let select_list = &after_select[..from_idx];
let mut seen: std::collections::HashSet<&str> = std::collections::HashSet::new();
for col in split_top_level_commas(select_list) {
let col = col.trim();
let alias = if let Some(idx) = col.rfind(" AS ") {
col[idx + " AS ".len()..].trim()
} else {
col
};
if !seen.insert(alias) {
return Err(crate::DjogiError::AliasCollision {
alias: alias.to_owned(),
});
}
}
Ok(())
}
fn split_top_level_commas(s: &str) -> Vec<&str> {
let mut out = Vec::new();
let mut depth = 0i32;
let mut start = 0usize;
let bytes = s.as_bytes();
for (i, &b) in bytes.iter().enumerate() {
match b {
b'(' => depth += 1,
b')' => depth -= 1,
b',' if depth == 0 => {
out.push(&s[start..i]);
start = i + 1;
}
_ => {}
}
}
out.push(&s[start..]);
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::descriptor::{ModelDescriptor, PkType};
use crate::query::condition::{Condition, FilterValue, Leaf, LookupOp};
use crate::query::queryset::QuerySet;
static FAKE_DESCRIPTOR: ModelDescriptor = ModelDescriptor {
type_name: "Fake",
table_name: "fakes",
pk_type: PkType::HeerIdDesc,
fields: &[],
partition_by: None,
has_outbox: false,
idempotency_key: None,
tenant_key: None,
cache_ttl: None,
rationale: None,
indexes: &[],
is_through: false,
fts: None,
app: None,
moved_from_app: None,
renamed_from: None,
exclusion_constraints: &[],
tree_edge: None,
proxy_for: None,
default_filter_sql: None,
computed_fields: &[],
table_comment: None,
storage_params: None,
tablespace: None,
};
struct Fake;
impl crate::model::__sealed::Sealed for Fake {}
#[allow(clippy::manual_async_fn)]
impl Model for Fake {
type Pk = i64;
type Fields = ();
fn table_name() -> &'static str {
"fakes"
}
fn pk_value(&self) -> &i64 {
unreachable!()
}
fn descriptor() -> &'static ModelDescriptor {
&FAKE_DESCRIPTOR
}
fn get(
_ctx: &mut crate::context::DjogiContext,
_id: i64,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn create(
_ctx: &mut crate::context::DjogiContext,
_v: Self,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn save<'ctx>(
&'ctx mut self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
fn delete(
self,
_ctx: &mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send {
async { unreachable!() }
}
fn refresh_from_db<'ctx>(
&'ctx self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
}
impl FromPgRow for Fake {
const COLUMNS: &'static [&'static str] = &["id"];
const COLUMN_LIST: &'static str = "id";
fn from_pg_row(_row: &tokio_postgres::Row) -> Result<Self, crate::DjogiError> {
unreachable!("SQL-text unit tests do not exercise row decode")
}
}
impl crate::pg::decode::FromJoinedPgRow for Fake {
fn from_joined_pg_row(
_row: &tokio_postgres::Row,
_prefix: &str,
) -> Result<Self, crate::DjogiError> {
unreachable!("SQL-text unit tests do not exercise row decode")
}
}
fn build_select<T: Model + FromPgRow>(qs: &QuerySet<T>) -> SqlAccumulator {
super::build_select(qs).expect("test predicate should lower to SQL")
}
fn build_select_joined<T: Model>(qs: &QuerySet<T>) -> SqlAccumulator {
super::build_select_joined(qs).expect("test predicate should lower to joined SQL")
}
fn build_count<T: Model>(qs: &QuerySet<T>) -> SqlAccumulator {
super::build_count(qs).expect("test predicate should lower to count SQL")
}
fn build_exists<T: Model>(qs: &QuerySet<T>) -> SqlAccumulator {
super::build_exists(qs).expect("test predicate should lower to exists SQL")
}
fn build_update<T: Model>(
qs: &QuerySet<T>,
assignments: &[crate::query::update::UpdateAssignment],
) -> SqlAccumulator {
super::build_update(qs, assignments).expect("test predicate should lower to update SQL")
}
fn build_delete<T: Model>(qs: &QuerySet<T>) -> SqlAccumulator {
super::build_delete(qs).expect("test predicate should lower to delete SQL")
}
fn build_grouped_annotated_select<T, K, A>(
gaq: &crate::query::grouped::GroupedAnnotatedQuerySet<T, K, A>,
) -> SqlAccumulator
where
T: Model,
K: crate::query::grouped::IntoGroupKeyTuple,
A: crate::query::annotate::IntoAggregateTuple,
{
super::build_grouped_annotated_select(gaq)
.expect("test predicate should lower to grouped annotated SQL")
}
#[test]
fn select_no_filter_omits_where() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_select(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "SELECT id FROM fakes");
}
#[test]
fn select_with_leaf_filter_emits_where_with_one_bind() {
let qs: QuerySet<Fake> =
QuerySet::new().filter(|_| Condition::Leaf(Leaf::eq_raw("a", FilterValue::Bool(true))));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE a = $1"), "got: {sql}");
}
#[test]
fn select_with_range_predicate_emits_postgres_range_operator() {
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| {
crate::query::field::FieldRef::<Fake, crate::Range<i32>>::new("span")
.adjacent_to(crate::Range::inclusive_exclusive(5_i32, 10_i32))
});
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE span -|- $1"), "got: {sql}");
}
#[test]
fn select_with_range_contains_element_emits_scalar_cast() {
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| {
crate::query::field::FieldRef::<Fake, crate::Range<i32>>::new("span").contains(3_i32)
});
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE span @> $1::int4"), "got: {sql}");
}
#[test]
fn select_with_and_uses_parentheses() {
let qs: QuerySet<Fake> = QuerySet::new()
.filter(|_| Condition::Leaf(Leaf::eq_raw("a", FilterValue::Bool(true))))
.filter(|_| Condition::Leaf(Leaf::eq_raw("b", FilterValue::Bool(false))));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE (a = $1 AND b = $2)"), "got: {sql}");
}
#[test]
fn select_with_exclude_wraps_not() {
let qs: QuerySet<Fake> = QuerySet::new()
.exclude(|_| Condition::Leaf(Leaf::eq_raw("a", FilterValue::Bool(true))));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE NOT (a = $1)"), "got: {sql}");
}
#[test]
fn select_distinct_plain_emits_distinct_keyword() {
let qs: QuerySet<Fake> = QuerySet::new().distinct();
let acc = build_select(&qs);
assert!(acc.sql().contains("SELECT DISTINCT id FROM fakes"));
}
#[test]
fn select_limit_offset_pushes_two_binds() {
let qs: QuerySet<Fake> = QuerySet::new().limit(10).offset(5);
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("LIMIT $1"), "got: {sql}");
assert!(sql.contains("OFFSET $2"), "got: {sql}");
}
#[test]
fn in_empty_list_renders_false() {
let leaf = Leaf::new("id", LookupOp::In, FilterValue::List(Vec::new()));
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| Condition::Leaf(leaf));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE FALSE"), "got: {sql}");
}
#[test]
fn not_in_empty_list_renders_true() {
let leaf = Leaf::new("id", LookupOp::NotIn, FilterValue::List(Vec::new()));
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| Condition::Leaf(leaf));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE TRUE"), "got: {sql}");
}
#[test]
fn in_list_emits_one_placeholder_per_element() {
let leaf = Leaf::new(
"id",
LookupOp::In,
FilterValue::List(vec![
FilterValue::I64(1),
FilterValue::I64(2),
FilterValue::I64(3),
]),
);
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| Condition::Leaf(leaf));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("id IN ($1, $2, $3)"), "got: {sql}");
}
#[test]
fn between_emits_two_binds() {
let leaf = Leaf::new(
"age",
LookupOp::Between,
FilterValue::Pair(
Box::new(FilterValue::I32(10)),
Box::new(FilterValue::I32(20)),
),
);
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| Condition::Leaf(leaf));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("age BETWEEN $1 AND $2"), "got: {sql}");
}
#[test]
fn is_null_takes_no_bind() {
let leaf = Leaf::new("deleted_at", LookupOp::IsNull, FilterValue::Null);
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| Condition::Leaf(leaf));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("deleted_at IS NULL"), "got: {sql}");
assert!(!sql.contains('$'), "expected no binds, got: {sql}");
}
#[test]
fn count_ignores_order_limit_offset() {
let qs: QuerySet<Fake> = QuerySet::new().limit(10).offset(5);
let acc = build_count(&qs);
let sql = acc.sql();
assert!(sql.starts_with("SELECT COUNT(*) FROM fakes"));
assert!(
!sql.contains("LIMIT"),
"count should not carry LIMIT: {sql}"
);
assert!(
!sql.contains("OFFSET"),
"count should not carry OFFSET: {sql}"
);
}
#[test]
fn exists_emits_limit_1_inside_subquery() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_exists(&qs);
let sql = acc.sql();
assert!(sql.contains("SELECT EXISTS(SELECT 1 FROM fakes"));
assert!(sql.contains("LIMIT 1"));
}
#[test]
fn order_by_asc_nulls_last_emits_expected_tokens() {
let qs: QuerySet<Fake> =
QuerySet::new().order_by(|_| crate::query::order::OrderExpr::Column {
column: "title",
direction: crate::query::order::Direction::Asc,
nulls: crate::query::order::NullsOrder::Last,
});
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("ORDER BY title ASC NULLS LAST"), "got: {sql}");
}
#[test]
fn distinct_on_emits_column_list() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.distinct = DistinctMode::On(vec!["title", "view_count"]);
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.contains("SELECT DISTINCT ON (title, view_count) id FROM fakes"),
"got: {sql}"
);
}
#[test]
fn like_escape_handles_percent_underscore_backslash() {
let leaf = Leaf::new(
"title",
LookupOp::IContains,
FilterValue::String("50% off_sale\\".to_string()),
);
let qs: QuerySet<Fake> = QuerySet::new().filter(|_| Condition::Leaf(leaf));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("title ILIKE $1"), "got: {sql}");
}
#[test]
fn escape_like_prefixes_special_chars() {
assert_eq!(escape_like("50%"), "50\\%");
assert_eq!(escape_like("a_b"), "a\\_b");
assert_eq!(escape_like("c\\d"), "c\\\\d");
assert_eq!(escape_like("plain"), "plain");
}
#[test]
fn count_with_distinct_plain_wraps_subquery() {
let qs: QuerySet<Fake> = QuerySet::new().distinct();
let acc = build_count(&qs);
let sql = acc.sql();
assert!(
sql.contains("SELECT COUNT(*) FROM (SELECT DISTINCT * FROM fakes)"),
"got: {sql}"
);
assert!(sql.contains(") AS sub"), "got: {sql}");
}
#[test]
fn count_with_distinct_on_wraps_subquery_with_order() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.distinct = DistinctMode::On(vec!["title", "view_count"]);
let acc = build_count(&qs);
let sql = acc.sql();
assert!(
sql.contains(
"SELECT COUNT(*) FROM (SELECT DISTINCT ON (title, view_count) * FROM fakes"
),
"got: {sql}"
);
assert!(sql.contains("ORDER BY title, view_count"), "got: {sql}");
assert!(sql.contains(") AS sub"), "got: {sql}");
}
#[test]
fn count_with_distinct_on_appends_user_ordering() {
let mut qs: QuerySet<Fake> =
QuerySet::new().order_by(|_| crate::query::order::OrderExpr::Column {
column: "view_count",
direction: crate::query::order::Direction::Desc,
nulls: crate::query::order::NullsOrder::Last,
});
qs.distinct = DistinctMode::On(vec!["title"]);
let acc = build_count(&qs);
let sql = acc.sql();
assert!(
sql.contains("ORDER BY title, view_count DESC NULLS LAST"),
"got: {sql}"
);
}
#[test]
fn count_without_distinct_omits_subquery() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_count(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "SELECT COUNT(*) FROM fakes");
}
#[test]
fn where_skipped_on_empty_and() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.condition = crate::query::Q::Condition(Condition::And(Vec::new()));
let acc = build_select(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "SELECT id FROM fakes");
}
#[test]
fn raw_sql_condition_wraps_in_parens() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.condition =
crate::query::Q::Condition(Condition::__from_raw_sql_fragment("active = TRUE"));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.contains("WHERE (active = TRUE)"),
"expected outer parens around proxy fragment, got: {sql}",
);
}
#[test]
fn raw_sql_condition_ands_with_user_filter() {
let raw = Condition::__from_raw_sql_fragment("active = TRUE");
let user = Condition::Leaf(Leaf::eq_raw("price", FilterValue::I64(100)));
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.condition = crate::query::Q::Condition(Condition::and(raw, user));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.contains("WHERE ((active = TRUE) AND price = $1)"),
"expected AND-composed WHERE clause, got: {sql}",
);
}
#[test]
fn where_skipped_on_nested_vacuous_and() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.condition = crate::query::Q::Condition(Condition::And(vec![
Condition::True,
Condition::And(Vec::new()),
]));
let acc = build_select(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "SELECT id FROM fakes");
}
#[test]
fn where_skipped_on_not_empty_or() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.condition =
crate::query::Q::Condition(Condition::Not(Box::new(Condition::Or(Vec::new()))));
let acc = build_select(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "SELECT id FROM fakes");
}
#[test]
fn update_single_assignment_emits_set_and_updated_at() {
use crate::query::update::{AssignmentValue, UpdateAssignment};
let a = UpdateAssignment {
column: "view_count",
value: AssignmentValue::Literal(FilterValue::I32(999)),
};
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_update(&qs, &[a]);
let sql = acc.sql();
assert!(
sql.contains("UPDATE fakes SET view_count = $1, updated_at = now()"),
"got: {sql}"
);
assert!(!sql.contains("WHERE"), "no filter -> no WHERE, got: {sql}");
}
#[test]
fn update_multiple_assignments_comma_separate_binds() {
use crate::query::update::{AssignmentValue, UpdateAssignment};
let a = UpdateAssignment {
column: "view_count",
value: AssignmentValue::Literal(FilterValue::I32(1)),
};
let b = UpdateAssignment {
column: "published",
value: AssignmentValue::Literal(FilterValue::Bool(true)),
};
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_update(&qs, &[a, b]);
let sql = acc.sql();
assert!(
sql.contains("SET view_count = $1, published = $2, updated_at = now()"),
"got: {sql}"
);
}
#[test]
fn update_with_filter_emits_where_with_bind_offset() {
use crate::query::update::{AssignmentValue, UpdateAssignment};
let a = UpdateAssignment {
column: "view_count",
value: AssignmentValue::Literal(FilterValue::I32(42)),
};
let qs: QuerySet<Fake> = QuerySet::new()
.filter(|_| Condition::Leaf(Leaf::eq_raw("published", FilterValue::Bool(true))));
let acc = build_update(&qs, &[a]);
let sql = acc.sql();
assert!(
sql.contains("SET view_count = $1, updated_at = now()"),
"got: {sql}"
);
assert!(sql.contains("WHERE published = $2"), "got: {sql}");
}
#[test]
fn delete_no_filter_emits_table_only() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_delete(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "DELETE FROM fakes");
}
#[test]
fn delete_with_filter_emits_where() {
let qs: QuerySet<Fake> = QuerySet::new()
.filter(|_| Condition::Leaf(Leaf::eq_raw("published", FilterValue::Bool(false))));
let acc = build_delete(&qs);
let sql = acc.sql();
assert!(sql.starts_with("DELETE FROM fakes"), "got: {sql}");
assert!(sql.contains("WHERE published = $1"), "got: {sql}");
}
#[test]
fn delete_vacuous_and_skips_where() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.condition = crate::query::Q::Condition(Condition::And(Vec::new()));
let acc = build_delete(&qs);
let sql = acc.sql().trim().to_string();
assert_eq!(sql, "DELETE FROM fakes");
}
struct FakeTarget;
impl crate::model::__sealed::Sealed for FakeTarget {}
#[allow(clippy::manual_async_fn)]
impl Model for FakeTarget {
type Pk = i64;
type Fields = ();
fn table_name() -> &'static str {
"fake_targets"
}
fn pk_value(&self) -> &i64 {
unreachable!()
}
fn descriptor() -> &'static ModelDescriptor {
unreachable!()
}
fn get(
_ctx: &mut crate::context::DjogiContext,
_id: i64,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn create(
_ctx: &mut crate::context::DjogiContext,
_v: Self,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn save<'ctx>(
&'ctx mut self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
fn delete(
self,
_ctx: &mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send {
async { unreachable!() }
}
fn refresh_from_db<'ctx>(
&'ctx self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
}
fn build_insert_select<S: Model, T: Model>(
qs: &QuerySet<S>,
columns: &[crate::query::insert_select::InsertSelectColumn<S, T>],
) -> SqlAccumulator {
super::build_insert_select::<S, T>(qs, columns)
.expect("test predicate should lower to insert-select SQL")
}
fn col_copy<S: Model, T: Model>(
target_column: &'static str,
source_column: &'static str,
) -> crate::query::insert_select::InsertSelectColumn<S, T> {
let target: crate::query::FieldRef<T, i32> = crate::query::FieldRef::new(target_column);
let source: crate::query::FieldRef<S, i32> = crate::query::FieldRef::new(source_column);
target.copy_from(source.as_insert_source())
}
#[test]
fn insert_select_no_filter_emits_bare_shape() {
let qs: QuerySet<Fake> = QuerySet::new();
let cols = vec![col_copy::<Fake, FakeTarget>("view_count", "score")];
let acc = build_insert_select::<Fake, FakeTarget>(&qs, &cols);
let sql = acc.sql().trim().to_string();
assert_eq!(
sql,
"INSERT INTO fake_targets (view_count) SELECT score FROM fakes"
);
}
#[test]
fn insert_select_multi_column_emits_lockstep_lists() {
let qs: QuerySet<Fake> = QuerySet::new();
let cols = vec![
col_copy::<Fake, FakeTarget>("a", "x"),
col_copy::<Fake, FakeTarget>("b", "y"),
col_copy::<Fake, FakeTarget>("c", "z"),
];
let acc = build_insert_select::<Fake, FakeTarget>(&qs, &cols);
let sql = acc.sql();
assert!(
sql.contains("INSERT INTO fake_targets (a, b, c) SELECT x, y, z FROM fakes"),
"got: {sql}"
);
}
#[test]
fn insert_select_with_filter_emits_where() {
let qs: QuerySet<Fake> = QuerySet::new()
.filter(|_| Condition::Leaf(Leaf::eq_raw("published", FilterValue::Bool(true))));
let cols = vec![col_copy::<Fake, FakeTarget>("view_count", "score")];
let acc = build_insert_select::<Fake, FakeTarget>(&qs, &cols);
let sql = acc.sql();
assert!(
sql.contains("FROM fakes WHERE published = $1"),
"got: {sql}"
);
}
#[test]
fn insert_select_with_literal_source_pushes_bind() {
let target: crate::query::FieldRef<FakeTarget, i32> =
crate::query::FieldRef::new("status_code");
let cols =
vec![target.copy_from(
crate::query::insert_select::InsertSelectSource::<Fake, _>::literal(7i32),
)];
let qs: QuerySet<Fake> = QuerySet::new()
.filter(|_| Condition::Leaf(Leaf::eq_raw("published", FilterValue::Bool(true))));
let acc = build_insert_select::<Fake, FakeTarget>(&qs, &cols);
let sql = acc.sql();
assert!(
sql.contains("INSERT INTO fake_targets (status_code) SELECT $1 FROM fakes"),
"got: {sql}"
);
assert!(sql.contains("WHERE published = $2"), "got: {sql}");
}
#[test]
fn insert_select_with_limit_offset_pushes_tail_binds() {
let qs: QuerySet<Fake> = QuerySet::new().limit(10).offset(5);
let cols = vec![col_copy::<Fake, FakeTarget>("view_count", "score")];
let acc = build_insert_select::<Fake, FakeTarget>(&qs, &cols);
let sql = acc.sql();
assert!(sql.contains("LIMIT $1"), "got: {sql}");
assert!(sql.contains("OFFSET $2"), "got: {sql}");
}
#[test]
fn insert_select_uses_source_table_in_from_not_target() {
let qs: QuerySet<Fake> = QuerySet::new();
let cols = vec![col_copy::<Fake, FakeTarget>("view_count", "score")];
let acc = build_insert_select::<Fake, FakeTarget>(&qs, &cols);
let sql = acc.sql();
assert!(sql.contains("INSERT INTO fake_targets"), "got: {sql}");
assert!(sql.contains("FROM fakes"), "got: {sql}");
assert!(
!sql.contains("FROM fake_targets"),
"INSERT...SELECT emitted FROM target table — sql: {sql}"
);
}
use crate::descriptor::{FieldDescriptor, FieldSqlType, field_descriptor, model_descriptor};
use crate::relation::select_related::ErasedSelectRelated;
static OWNERS_JOIN_DESC: ModelDescriptor = ModelDescriptor {
..model_descriptor(
"Owner",
"owners_p3",
PkType::HeerId,
&[FieldDescriptor {
unique: true,
indexed: true,
..field_descriptor("id", FieldSqlType::BigInt, false)
}],
)
};
fn owners_join_descriptor() -> &'static ModelDescriptor {
&OWNERS_JOIN_DESC
}
fn dummy_join_decoder(
_row: &tokio_postgres::Row,
_prefix: &str,
) -> Result<Option<Box<dyn std::any::Any + Send + Sync>>, crate::DjogiError> {
unreachable!("dummy decoder should not run in SQL-emission tests")
}
fn owner_path() -> ErasedSelectRelated {
ErasedSelectRelated {
source_column: "owner_id",
child_table: "owners_p3",
decoder: dummy_join_decoder,
child_descriptor: owners_join_descriptor,
}
}
#[test]
fn joined_select_qualifies_where_column_refs_with_parent_table() {
let mut qs: QuerySet<Fake> =
QuerySet::new().filter(|_| Condition::Leaf(Leaf::eq_raw("id", FilterValue::I64(42))));
qs.select_related_paths.push(owner_path());
let acc = build_select_joined(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE fakes.id = $1"), "got: {sql}");
assert!(
sql.contains("LEFT JOIN owners_p3 rel_owner_id"),
"LEFT JOIN missing, got: {sql}"
);
}
#[test]
fn joined_select_qualifies_order_by_column_refs() {
let mut qs: QuerySet<Fake> =
QuerySet::new().order_by(|_| crate::query::order::OrderExpr::Column {
column: "created_at",
direction: crate::query::order::Direction::Asc,
nulls: crate::query::order::NullsOrder::Default,
});
qs.select_related_paths.push(owner_path());
let acc = build_select_joined(&qs);
let sql = acc.sql();
assert!(sql.contains("ORDER BY fakes.created_at ASC"), "got: {sql}");
}
#[test]
fn joined_select_qualifies_distinct_on_column_refs() {
let mut qs: QuerySet<Fake> = QuerySet::new();
qs.distinct = DistinctMode::On(vec!["id"]);
qs.select_related_paths.push(owner_path());
let acc = build_select_joined(&qs);
let sql = acc.sql();
assert!(sql.contains("SELECT DISTINCT ON (fakes.id)"), "got: {sql}");
}
#[test]
fn non_joined_select_leaves_column_refs_bare() {
let qs: QuerySet<Fake> =
QuerySet::new().filter(|_| Condition::Leaf(Leaf::eq_raw("id", FilterValue::I64(42))));
let acc = build_select(&qs);
let sql = acc.sql();
assert!(sql.contains("WHERE id = $1"), "got: {sql}");
assert!(
!sql.contains("fakes.id"),
"bare query must not qualify: {sql}"
);
}
#[test]
fn select_for_update_appends_lock_tail() {
let qs: QuerySet<Fake> = QuerySet::new().select_for_update();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE"),
"expected FOR UPDATE tail, got: {sql}"
);
assert!(
!sql.contains("NOWAIT") && !sql.contains("SKIP LOCKED"),
"select_for_update must not escalate to NOWAIT / SKIP LOCKED"
);
}
#[test]
fn nowait_appends_for_update_nowait_tail() {
let qs: QuerySet<Fake> = QuerySet::new().nowait();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE NOWAIT"),
"expected FOR UPDATE NOWAIT tail, got: {sql}"
);
}
#[test]
fn skip_locked_appends_for_update_skip_locked_tail() {
let qs: QuerySet<Fake> = QuerySet::new().skip_locked();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE SKIP LOCKED"),
"expected FOR UPDATE SKIP LOCKED tail, got: {sql}"
);
}
#[test]
fn lock_tail_follows_limit_and_offset() {
let qs: QuerySet<Fake> = QuerySet::new().limit(10).offset(5).select_for_update();
let acc = build_select(&qs);
let sql = acc.sql();
let limit_idx = sql.find("LIMIT").expect("LIMIT must appear");
let offset_idx = sql.find("OFFSET").expect("OFFSET must appear");
let lock_idx = sql.find("FOR UPDATE").expect("FOR UPDATE must appear");
assert!(
limit_idx < offset_idx && offset_idx < lock_idx,
"expected LIMIT ... OFFSET ... FOR UPDATE order, got: {sql}"
);
}
#[test]
fn lock_builder_last_call_wins_across_nowait_skip_locked() {
let qs: QuerySet<Fake> = QuerySet::new().nowait().skip_locked();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE SKIP LOCKED"),
"expected skip_locked to win over nowait, got: {sql}"
);
}
#[test]
fn select_for_share_appends_lock_tail() {
let qs: QuerySet<Fake> = QuerySet::new().select_for_share();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR SHARE"),
"expected FOR SHARE tail, got: {sql}"
);
assert!(
!sql.contains("NOWAIT") && !sql.contains("SKIP LOCKED"),
"select_for_share must not escalate to NOWAIT / SKIP LOCKED"
);
assert!(
!sql.contains("FOR UPDATE"),
"select_for_share must not emit FOR UPDATE: {sql}"
);
}
#[test]
fn for_share_nowait_appends_for_share_nowait_tail() {
let qs: QuerySet<Fake> = QuerySet::new().for_share_nowait();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR SHARE NOWAIT"),
"expected FOR SHARE NOWAIT tail, got: {sql}"
);
assert!(
!sql.contains("FOR UPDATE"),
"for_share_nowait must not emit FOR UPDATE: {sql}"
);
}
#[test]
fn for_share_skip_locked_appends_for_share_skip_locked_tail() {
let qs: QuerySet<Fake> = QuerySet::new().for_share_skip_locked();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR SHARE SKIP LOCKED"),
"expected FOR SHARE SKIP LOCKED tail, got: {sql}"
);
assert!(
!sql.contains("FOR UPDATE"),
"for_share_skip_locked must not emit FOR UPDATE: {sql}"
);
}
#[test]
fn for_share_tail_follows_limit_and_offset() {
let qs: QuerySet<Fake> = QuerySet::new().limit(10).offset(5).select_for_share();
let acc = build_select(&qs);
let sql = acc.sql();
let limit_idx = sql.find("LIMIT").expect("LIMIT must appear");
let offset_idx = sql.find("OFFSET").expect("OFFSET must appear");
let lock_idx = sql.find("FOR SHARE").expect("FOR SHARE must appear");
assert!(
limit_idx < offset_idx && offset_idx < lock_idx,
"expected LIMIT ... OFFSET ... FOR SHARE order, got: {sql}"
);
}
#[test]
fn for_share_builders_last_call_wins() {
let qs: QuerySet<Fake> = QuerySet::new().for_share_nowait().for_share_skip_locked();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR SHARE SKIP LOCKED"),
"expected for_share_skip_locked to win, got: {sql}"
);
}
#[test]
fn for_share_then_for_update_last_call_wins() {
let qs: QuerySet<Fake> = QuerySet::new().select_for_share().select_for_update();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE"),
"expected last-call-wins flip to FOR UPDATE, got: {sql}"
);
assert!(
!sql.contains("FOR SHARE"),
"select_for_update after select_for_share must clear the FOR SHARE tail: {sql}"
);
}
#[test]
fn nowait_after_select_for_share_promotes_to_for_update_nowait() {
let qs: QuerySet<Fake> = QuerySet::new().select_for_share().nowait();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE NOWAIT"),
"expected .nowait() after .select_for_share() to promote to FOR UPDATE NOWAIT, got: {sql}"
);
assert!(
!sql.contains("FOR SHARE"),
".nowait() must replace the FOR SHARE tail outright (documented footgun): {sql}"
);
}
#[test]
fn skip_locked_after_select_for_share_promotes_to_for_update_skip_locked() {
let qs: QuerySet<Fake> = QuerySet::new().select_for_share().skip_locked();
let acc = build_select(&qs);
let sql = acc.sql();
assert!(
sql.trim_end().ends_with("FOR UPDATE SKIP LOCKED"),
"expected .skip_locked() after .select_for_share() to promote to FOR UPDATE SKIP LOCKED, got: {sql}"
);
assert!(
!sql.contains("FOR SHARE"),
".skip_locked() must replace the FOR SHARE tail outright (documented footgun): {sql}"
);
}
#[test]
fn build_grouped_annotated_select_emits_grouping_sets() {
use crate::expr::AggregateExpr;
use crate::query::field::FieldRef;
use crate::query::grouped::{GroupedAnnotatedQuerySet, GroupingMode};
use std::marker::PhantomData;
let qs: QuerySet<Fake> = QuerySet::new();
let vals: FieldRef<Fake, i64> = FieldRef::new("amount");
let gaq: GroupedAnnotatedQuerySet<Fake, (), AggregateExpr<i64>> = {
let gq = crate::query::grouped::GroupedQuerySet {
qs,
keys: (),
grouping: GroupingMode::Sets(vec![vec!["org_id"], vec!["region"]]),
#[cfg(feature = "spatial")]
spatial_source: None,
_k: PhantomData,
};
gq.annotate(|_| vals.sum())
};
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
!sql.contains("SELECT ,"),
"unit-key grouping sets must not emit a leading SELECT comma: {sql}"
);
assert!(
sql.starts_with("SELECT (SUM(amount))::BIGINT AS __djogi_agg_0 FROM fakes AS t"),
"unit-key grouping sets must start with the aggregate projection, got: {sql}"
);
assert!(
sql.contains("GROUPING SETS ((org_id), (region))"),
"expected GROUPING SETS clause, got: {sql}"
);
}
#[test]
fn build_grouped_annotated_select_emits_rollup() {
use crate::expr::AggregateExpr;
use crate::query::field::FieldRef;
use crate::query::grouped::{GroupedAnnotatedQuerySet, GroupingMode};
use std::marker::PhantomData;
let qs: QuerySet<Fake> = QuerySet::new();
let keys: FieldRef<Fake, i64> = FieldRef::new("org_id");
let vals: FieldRef<Fake, i64> = FieldRef::new("amount");
let gaq: GroupedAnnotatedQuerySet<Fake, FieldRef<Fake, i64>, AggregateExpr<i64>> = {
let gq = crate::query::grouped::GroupedQuerySet {
qs,
keys,
grouping: GroupingMode::Rollup,
#[cfg(feature = "spatial")]
spatial_source: None,
_k: PhantomData,
};
gq.annotate(|_| vals.sum())
};
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("GROUP BY ROLLUP (org_id)"),
"expected ROLLUP clause, got: {sql}"
);
}
#[test]
fn build_grouped_annotated_select_emits_cube() {
use crate::expr::AggregateExpr;
use crate::query::field::FieldRef;
use crate::query::grouped::{GroupedAnnotatedQuerySet, GroupingMode};
use std::marker::PhantomData;
let qs: QuerySet<Fake> = QuerySet::new();
let keys: FieldRef<Fake, i64> = FieldRef::new("org_id");
let vals: FieldRef<Fake, i64> = FieldRef::new("amount");
let gaq: GroupedAnnotatedQuerySet<Fake, FieldRef<Fake, i64>, AggregateExpr<i64>> = {
let gq = crate::query::grouped::GroupedQuerySet {
qs,
keys,
grouping: GroupingMode::Cube,
#[cfg(feature = "spatial")]
spatial_source: None,
_k: PhantomData,
};
gq.annotate(|_| vals.sum())
};
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("GROUP BY CUBE (org_id)"),
"expected CUBE clause, got: {sql}"
);
}
#[test]
fn alias_collision_detected_in_grouped_select() {
let ok_sql = "SELECT org_id, SUM(amount) AS __djogi_agg_0 FROM txns GROUP BY org_id";
let result = assert_no_alias_collision(ok_sql);
assert!(result.is_ok(), "expected no collision, got: {:?}", result);
}
#[test]
fn alias_collision_names_both_columns_in_error() {
let bad_sql = "SELECT foo AS dup, bar AS dup FROM t";
let err = assert_no_alias_collision(bad_sql).unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("dup"),
"error should name the conflicting alias 'dup', got: {msg}"
);
}
#[test]
fn alias_collision_bare_name_collision_detected() {
let bad_sql = "SELECT foo, foo FROM t";
let err = assert_no_alias_collision(bad_sql).unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("foo"),
"error should name the conflicting alias 'foo', got: {msg}"
);
}
#[test]
fn alias_collision_mixed_as_and_bare_collision_detected() {
let bad_sql = "SELECT org_id, SUM(amount) AS org_id FROM t GROUP BY org_id";
let err = assert_no_alias_collision(bad_sql).unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("org_id"),
"error should name the conflicting alias 'org_id', got: {msg}"
);
}
#[test]
fn alias_collision_happy_path_grouped_queryset() {
use crate::expr::AggregateExpr;
use crate::query::field::FieldRef;
let qs: QuerySet<Fake> = QuerySet::new();
let keys: FieldRef<Fake, i64> = FieldRef::new("org_id");
let vals: FieldRef<Fake, i64> = FieldRef::new("amount");
let gaq = qs.group_by(|_| keys).annotate(|_| vals.sum());
let acc =
build_grouped_annotated_select::<Fake, FieldRef<Fake, i64>, AggregateExpr<i64>>(&gaq);
let result = assert_no_alias_collision(acc.sql());
assert!(
result.is_ok(),
"expected no alias collision in grouped queryset, got: {:?}",
result
);
}
#[cfg(feature = "spatial")]
struct FakeRegion;
#[cfg(feature = "spatial")]
impl crate::model::__sealed::Sealed for FakeRegion {}
#[cfg(feature = "spatial")]
#[allow(clippy::manual_async_fn)]
impl Model for FakeRegion {
type Pk = i64;
type Fields = ();
fn table_name() -> &'static str {
"neighborhoods"
}
fn pk_value(&self) -> &i64 {
unreachable!()
}
fn descriptor() -> &'static ModelDescriptor {
unreachable!()
}
fn get(
_ctx: &mut crate::context::DjogiContext,
_id: i64,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn create(
_ctx: &mut crate::context::DjogiContext,
_v: Self,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn save<'ctx>(
&'ctx mut self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
fn delete(
self,
_ctx: &mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send {
async { unreachable!() }
}
fn refresh_from_db<'ctx>(
&'ctx self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
}
#[cfg(feature = "spatial")]
fn make_spatial_gaq(
spec: crate::query::spatial_grouping::SpatialJoinSpec,
) -> crate::query::grouped::GroupedAnnotatedQuerySet<
Fake,
crate::query::spatial_grouping::RegionKey<FakeRegion>,
crate::expr::AggregateExpr<i64>,
> {
use std::marker::PhantomData;
let keys = crate::query::spatial_grouping::RegionKey::<FakeRegion> {
region_pk: None,
r_pk_col: Some(spec.r_pk_col),
_phantom: PhantomData,
};
let agg: crate::expr::AggregateExpr<i64> =
crate::query::field::FieldRef::<Fake, i64>::new("id").count_star();
crate::query::grouped::GroupedAnnotatedQuerySet {
qs: QuerySet::new(),
keys,
grouping: crate::query::grouped::GroupingMode::Plain,
aggregates: agg,
having: None,
order: Vec::new(),
limit: None,
offset: None,
spatial_source: Some(crate::query::grouped::SpatialGroupSource::Join(spec)),
_k: PhantomData,
_a: PhantomData,
}
}
#[cfg(feature = "spatial")]
#[test]
fn spatial_join_emits_left_join_with_st_covers() {
let spec = crate::query::spatial_grouping::SpatialJoinSpec {
t_geo_col: "location",
r_table: "neighborhoods",
r_geo_col: "boundary",
r_pk_col: "id",
};
let gaq = make_spatial_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("LEFT JOIN neighborhoods AS r ON ST_Covers(r.boundary, t.location)"),
"expected LEFT JOIN with ST_Covers (geography overload), got: {sql}"
);
assert!(
!sql.contains("ST_Contains"),
"must not emit ST_Contains (no geography overload in PostGIS 3.x); got: {sql}"
);
}
#[cfg(feature = "spatial")]
#[test]
fn spatial_join_groups_by_region_pk_qualified() {
let spec = crate::query::spatial_grouping::SpatialJoinSpec {
t_geo_col: "location",
r_table: "neighborhoods",
r_geo_col: "boundary",
r_pk_col: "id",
};
let gaq = make_spatial_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("GROUP BY r.id"),
"expected GROUP BY r.id, got: {sql}"
);
}
#[cfg(feature = "spatial")]
#[test]
fn spatial_join_select_list_starts_with_region_pk_alias() {
let spec = crate::query::spatial_grouping::SpatialJoinSpec {
t_geo_col: "location",
r_table: "neighborhoods",
r_geo_col: "boundary",
r_pk_col: "id",
};
let gaq = make_spatial_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.starts_with("SELECT r.id AS rk0"),
"expected SELECT to start with 'SELECT r.id AS rk0', got: {sql}"
);
}
#[cfg(feature = "spatial")]
#[test]
fn spatial_join_clause_order_is_correct() {
let spec = crate::query::spatial_grouping::SpatialJoinSpec {
t_geo_col: "location",
r_table: "neighborhoods",
r_geo_col: "boundary",
r_pk_col: "id",
};
let gaq = make_spatial_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
let from_pos = sql.find("FROM fakes AS t").unwrap();
let join_pos = sql.find("LEFT JOIN neighborhoods").unwrap();
let group_pos = sql.find("GROUP BY").unwrap();
assert!(
from_pos < join_pos,
"FROM must precede LEFT JOIN; got: {sql}"
);
assert!(
join_pos < group_pos,
"LEFT JOIN must precede GROUP BY; got: {sql}"
);
}
#[cfg(feature = "spatial")]
fn make_cluster_gaq(
spec: crate::query::spatial_grouping::ClusterSpec,
) -> crate::query::grouped::GroupedAnnotatedQuerySet<
Fake,
crate::query::spatial_grouping::ClusterId,
crate::expr::AggregateExpr<i64>,
> {
use std::marker::PhantomData;
let agg: crate::expr::AggregateExpr<i64> =
crate::query::field::FieldRef::<Fake, i64>::new("id").count_star();
crate::query::grouped::GroupedAnnotatedQuerySet {
qs: QuerySet::new(),
keys: crate::query::spatial_grouping::ClusterId(None),
grouping: crate::query::grouped::GroupingMode::Plain,
aggregates: agg,
having: None,
order: Vec::new(),
limit: None,
offset: None,
spatial_source: Some(crate::query::grouped::SpatialGroupSource::Cluster(spec)),
_k: PhantomData,
_a: PhantomData,
}
}
#[cfg(feature = "spatial")]
#[test]
fn cluster_grouped_select_emits_st_cluster_dbscan_with_geometry_cast() {
use crate::query::spatial_grouping::ClusterSpec;
let spec = ClusterSpec {
t_geo_col: "location",
eps_degrees: 0.004491,
minpoints: 3,
};
let gaq = make_cluster_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("ST_ClusterDBSCAN(t.location::geometry,"),
"expected ST_ClusterDBSCAN with ::geometry cast, got: {sql}"
);
assert!(
sql.contains("OVER () AS cluster_id"),
"expected OVER () AS cluster_id, got: {sql}"
);
assert!(
sql.contains("GROUP BY cluster_id"),
"expected GROUP BY cluster_id, got: {sql}"
);
}
#[cfg(feature = "spatial")]
#[test]
fn cluster_grouped_select_binds_eps_and_minpoints() {
use crate::query::spatial_grouping::ClusterSpec;
let spec = ClusterSpec {
t_geo_col: "location",
eps_degrees: 0.00449,
minpoints: 5,
};
let gaq = make_cluster_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("$1") && sql.contains("$2"),
"expected $1 (eps) and $2 (minpoints) bind slots, got: {sql}"
);
assert!(
!sql.contains("LEFT JOIN"),
"cluster path should not emit LEFT JOIN, got: {sql}"
);
}
#[cfg(feature = "spatial")]
#[test]
fn cluster_grouped_select_wraps_window_in_subquery() {
use crate::query::spatial_grouping::ClusterSpec;
let spec = ClusterSpec {
t_geo_col: "location",
eps_degrees: 0.004491,
minpoints: 3,
};
let gaq = make_cluster_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.starts_with("SELECT cluster_id"),
"outer SELECT must start with 'SELECT cluster_id', got: {sql}"
);
assert!(
sql.contains("FROM (SELECT t.*, ST_ClusterDBSCAN("),
"window call must be wrapped in an inner subquery; got: {sql}"
);
assert!(
sql.contains(") AS t GROUP BY cluster_id"),
"subquery must be aliased 'AS t' and outer must GROUP BY cluster_id; got: {sql}"
);
assert!(
!sql.starts_with("SELECT ST_ClusterDBSCAN"),
"outer SELECT must not begin with the inline window form; got: {sql}"
);
}
#[cfg(feature = "spatial")]
fn make_geohash_gaq(
spec: crate::query::spatial_grouping::GeohashSpec,
) -> crate::query::grouped::GroupedAnnotatedQuerySet<
Fake,
crate::query::spatial_grouping::GeohashKey,
crate::expr::AggregateExpr<i64>,
> {
use std::marker::PhantomData;
let agg: crate::expr::AggregateExpr<i64> =
crate::query::field::FieldRef::<Fake, i64>::new("id").count_star();
crate::query::grouped::GroupedAnnotatedQuerySet {
qs: QuerySet::new(),
keys: crate::query::spatial_grouping::GeohashKey(None),
grouping: crate::query::grouped::GroupingMode::Plain,
aggregates: agg,
having: None,
order: Vec::new(),
limit: None,
offset: None,
spatial_source: Some(crate::query::grouped::SpatialGroupSource::Geohash(spec)),
_k: PhantomData,
_a: PhantomData,
}
}
#[cfg(feature = "spatial")]
#[test]
fn geohash_grouped_select_emits_st_geohash_with_geometry_cast() {
use crate::query::spatial_grouping::GeohashSpec;
let spec = GeohashSpec {
t_geo_col: "location",
precision: 5,
};
let gaq = make_geohash_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("ST_GeoHash(t.location::geometry,"),
"expected ST_GeoHash with ::geometry cast, got: {sql}"
);
assert!(
sql.contains("AS geohash"),
"expected AS geohash alias, got: {sql}"
);
assert!(
sql.contains("GROUP BY geohash"),
"expected GROUP BY geohash, got: {sql}"
);
}
#[cfg(feature = "spatial")]
#[test]
fn geohash_grouped_select_binds_precision_only() {
use crate::query::spatial_grouping::GeohashSpec;
let spec = GeohashSpec {
t_geo_col: "location",
precision: 7,
};
let gaq = make_geohash_gaq(spec);
let acc = build_grouped_annotated_select(&gaq);
let sql = acc.sql();
assert!(
sql.contains("$1"),
"expected $1 (precision) bind slot, got: {sql}"
);
assert!(
!sql.contains("LEFT JOIN"),
"geohash path should not emit LEFT JOIN, got: {sql}"
);
}
fn build_update_returning_pairs<T: Model + FromPgRow + crate::pg::decode::FromJoinedPgRow>(
qs: &QuerySet<T>,
assignments: &[crate::query::update::UpdateAssignment],
) -> SqlAccumulator {
super::build_update_returning_pairs(qs, assignments)
.expect("update_returning_pairs should build successfully")
}
fn build_delete_returning<T: Model + FromPgRow + crate::pg::decode::FromJoinedPgRow>(
qs: &QuerySet<T>,
) -> SqlAccumulator {
super::build_delete_returning(qs).expect("build_delete_returning should build successfully")
}
#[test]
fn update_returning_pairs_emits_returning_with_old_and_new_clause() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(999i32));
let acc = build_update_returning_pairs(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
assert!(
sql.contains("RETURNING WITH (OLD AS __djogi_old, NEW AS __djogi_new)"),
"expected RETURNING WITH clause, got: {sql}"
);
}
#[test]
fn update_returning_pairs_includes_old_id_alias() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(999i32));
let acc = build_update_returning_pairs(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
assert!(sql.contains("\"o0\""), "expected o0 alias, got: {sql}");
}
#[test]
fn update_returning_pairs_includes_new_id_alias() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(999i32));
let acc = build_update_returning_pairs(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
assert!(sql.contains("\"n0\""), "expected n0 alias, got: {sql}");
}
#[test]
fn update_returning_pairs_bind_order_assignments_before_filter() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(999i32));
let acc = build_update_returning_pairs(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
let returning_pos = sql.find("RETURNING").expect("should contain RETURNING");
let bind_pos = sql.find("$1").expect("should contain $1");
assert!(
bind_pos < returning_pos,
"bind slot $1 should appear before RETURNING clause, got: {sql}"
);
}
#[test]
fn delete_returning_emits_returning_with_old_clause() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_delete_returning(&qs);
let sql = acc.sql();
assert!(
sql.contains("RETURNING WITH (OLD AS __djogi_old)"),
"expected DELETE RETURNING WITH OLD clause, got: {sql}"
);
}
#[test]
fn delete_returning_includes_old_id_alias() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_delete_returning(&qs);
let sql = acc.sql();
assert!(
sql.contains("\"o0\""),
"expected o0 alias in DELETE returning, got: {sql}"
);
}
#[test]
fn delete_returning_does_not_include_new_projection() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_delete_returning(&qs);
let sql = acc.sql();
assert!(
!sql.contains("__djogi_new"),
"DELETE returning must not include new projection, got: {sql}"
);
}
#[test]
fn update_returning_pairs_projection_includes_both_sides_shape() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(0i32));
let acc = build_update_returning_pairs(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
assert!(sql.contains("OLD AS __djogi_old"), "{sql}");
assert!(sql.contains("NEW AS __djogi_new"), "{sql}");
assert!(sql.contains("\"o0\""), "{sql}");
assert!(sql.contains("\"n0\""), "{sql}");
}
#[test]
fn delete_returning_projection_is_old_only_shape() {
let qs: QuerySet<Fake> = QuerySet::new();
let acc = build_delete_returning(&qs);
let sql = acc.sql();
assert!(sql.contains("OLD AS __djogi_old"), "{sql}");
assert!(!sql.contains("NEW"), "{sql}");
assert!(sql.contains("\"o0\""), "{sql}");
}
fn build_update_returning_ids<T: Model + FromPgRow>(
qs: &QuerySet<T>,
assignments: &[crate::query::update::UpdateAssignment],
) -> SqlAccumulator {
super::build_update_returning_ids(qs, assignments)
.expect("update_returning_ids should build successfully")
}
#[test]
fn update_returning_ids_emits_returning_pk_clause() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(999i32));
let acc = build_update_returning_ids(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
assert!(
sql.contains("RETURNING id"),
"expected RETURNING id clause for bulk update cache invalidation, got: {sql}"
);
}
#[test]
fn update_returning_ids_bind_order_assignments_before_filter() {
let f: crate::query::field::FieldRef<Fake, i32> =
crate::query::field::FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new()
.filter(|_| Condition::Leaf(Leaf::eq_raw("published", FilterValue::Bool(true))));
let stmt = qs.update(|_| f.set(999i32));
let acc = build_update_returning_ids(&stmt.qs, &stmt.assignments);
let sql = acc.sql();
let set_pos = sql.find("SET").expect("should contain SET");
let where_pos = sql.find("WHERE").expect("should contain WHERE");
let returning_pos = sql.find("RETURNING").expect("should contain RETURNING");
assert!(
set_pos < where_pos && where_pos < returning_pos,
"expected SET ... WHERE ... RETURNING order, got: {sql}"
);
}
}