use crate::sql::{Dialect, SqlValue};
use crate::tenancy::ResolvedScope;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Agg {
Count,
Sum,
Avg,
Min,
Max,
}
impl Agg {
fn keyword(self) -> &'static str {
match self {
Self::Count => "count",
Self::Sum => "sum",
Self::Avg => "avg",
Self::Min => "min",
Self::Max => "max",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BinOp {
Add,
Sub,
Mul,
Div,
Mod,
}
impl BinOp {
fn symbol(self) -> &'static str {
match self {
Self::Add => "+",
Self::Sub => "-",
Self::Mul => "*",
Self::Div => "/",
Self::Mod => "%",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Func {
Lower,
Upper,
Length,
Trim,
Abs,
Round,
Coalesce,
Now,
}
impl Func {
fn spec(self) -> (&'static str, usize, Option<usize>) {
match self {
Self::Lower => ("lower", 1, Some(1)),
Self::Upper => ("upper", 1, Some(1)),
Self::Length => ("length", 1, Some(1)),
Self::Trim => ("trim", 1, Some(1)),
Self::Abs => ("abs", 1, Some(1)),
Self::Round => ("round", 1, Some(2)),
Self::Coalesce => ("coalesce", 2, None),
Self::Now => ("current_timestamp", 0, Some(0)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Metric {
Cosine,
L2,
}
impl Metric {
fn operator(self) -> &'static str {
match self {
Self::Cosine => "<=>",
Self::L2 => "<->",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RelArg {
Star,
Column(String),
}
#[derive(Debug, Clone, PartialEq)]
pub enum Expr {
Column(String),
Value(SqlValue),
Star,
Aggregate(Agg, Box<Self>),
Binary(BinOp, Box<Self>, Box<Self>),
Func(Func, Vec<Self>),
JsonExtract(Box<Self>, Vec<String>),
Distance {
left: Box<Self>,
right: Box<Self>,
metric: Metric,
},
VectorLiteral(String),
RelatedAggregate {
agg: Agg,
arg: RelArg,
table: String,
filter: Box<Predicate>,
},
Case {
branches: Vec<(Predicate, Self)>,
otherwise: Option<Box<Self>>,
},
JsonExtractDyn(Box<Self>, Box<Self>),
JsonConcat(Box<Self>, Box<Self>),
RelatedScalar {
column: String,
table: String,
filter: Box<Predicate>,
},
IsOwn,
}
impl Expr {
pub fn col(name: impl Into<String>) -> Self {
Self::Column(name.into())
}
pub fn val(v: impl Into<SqlValue>) -> Self {
Self::Value(v.into())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CmpOp {
Eq,
Ne,
Lt,
Le,
Gt,
Ge,
}
impl CmpOp {
fn symbol(self) -> &'static str {
match self {
Self::Eq => "=",
Self::Ne => "<>",
Self::Lt => "<",
Self::Le => "<=",
Self::Gt => ">",
Self::Ge => ">=",
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Predicate {
And(Vec<Self>),
Or(Vec<Self>),
Not(Box<Self>),
Cmp { left: Expr, op: CmpOp, right: Expr },
Between {
expr: Expr,
low: Expr,
high: Expr,
negated: bool,
},
In {
expr: Expr,
values: Vec<Expr>,
negated: bool,
},
Like {
expr: Expr,
pattern: String,
insensitive: bool,
negated: bool,
},
Null { expr: Expr, negated: bool },
InSubquery {
expr: Expr,
column: String,
table: String,
filter: Box<Self>,
negated: bool,
},
}
pub fn all(preds: impl IntoIterator<Item = Predicate>) -> Predicate {
Predicate::And(preds.into_iter().collect())
}
pub fn any(preds: impl IntoIterator<Item = Predicate>) -> Predicate {
Predicate::Or(preds.into_iter().collect())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum JoinKind {
Inner,
Left,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Join {
pub kind: JoinKind,
pub table: String,
pub alias: Option<String>,
pub on: Predicate,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Direction {
Asc,
Desc,
}
#[derive(Debug, Clone, PartialEq)]
pub struct OrderBy {
pub expr: Expr,
pub dir: Direction,
}
#[derive(Debug, Clone, PartialEq)]
pub struct SelectItem {
pub expr: Expr,
pub alias: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ScopeMode {
#[default]
Own,
OwnOrNull,
NullOnly,
All,
}
#[derive(Debug, Clone, PartialEq, Default)]
pub enum TableKeys {
#[default]
Uniform,
PerTable(std::collections::BTreeMap<String, ResolvedScope>),
PerTableTarget {
keys: std::collections::BTreeMap<String, ResolvedScope>,
public: std::collections::BTreeMap<String, Vec<PublicTermSql>>,
write: std::collections::BTreeSet<String>,
require_public: bool,
},
}
#[derive(Debug, Clone, PartialEq)]
pub enum PublicTermSql {
Cmp {
column: String,
op: CmpOp,
value: SqlValue,
},
Null { column: String, negated: bool },
}
pub fn lower_public_terms(pred: &crate::tenancy::PublicPredicate) -> Vec<PublicTermSql> {
use crate::tenancy::{PublicCmp, PublicLiteral, PublicTerm};
pred.terms
.iter()
.map(|t| match t {
PublicTerm::Cmp { column, op, value } => {
let op = match op {
PublicCmp::Eq => CmpOp::Eq,
PublicCmp::Ne => CmpOp::Ne,
PublicCmp::Lt => CmpOp::Lt,
PublicCmp::Le => CmpOp::Le,
PublicCmp::Gt => CmpOp::Gt,
PublicCmp::Ge => CmpOp::Ge,
};
let value = match value {
PublicLiteral::Bool(b) => SqlValue::Boolean(*b),
PublicLiteral::Int(n) => SqlValue::Integer(*n),
PublicLiteral::Text(s) => SqlValue::Text(s.clone()),
};
PublicTermSql::Cmp {
column: column.clone(),
op,
value,
}
}
PublicTerm::Null { column, negated } => PublicTermSql::Null {
column: column.clone(),
negated: *negated,
},
})
.collect()
}
#[derive(Debug, Clone, PartialEq)]
pub struct Scope {
pub column: String,
pub value: Option<SqlValue>,
pub session: Option<SqlValue>,
pub mode: ScopeMode,
pub keys: TableKeys,
}
impl Scope {
fn resolve_table(&self, table: &str) -> Result<ResolvedScope, OrmError> {
match &self.keys {
TableKeys::Uniform => Ok(ResolvedScope::Column(self.column.clone())),
TableKeys::PerTable(m) | TableKeys::PerTableTarget { keys: m, .. } => m
.get(table)
.cloned()
.ok_or_else(|| OrmError::TenancyUndeclared(table.to_string())),
}
}
fn public_pred(
&self,
table: &str,
qualifier: Option<&str>,
) -> Result<Option<Predicate>, OrmError> {
let TableKeys::PerTableTarget {
public,
require_public,
..
} = &self.keys
else {
return Ok(None);
};
let terms = match public.get(table) {
Some(t) => t,
None if !require_public => return Ok(None),
None => return Err(OrmError::PublicSubsetUndeclared(table.to_string())),
};
let mut preds = Vec::with_capacity(terms.len());
for term in terms {
match term {
PublicTermSql::Cmp { column, op, value } => {
ident(column)?;
preds.push(Predicate::Cmp {
left: Self::col_expr(column, qualifier),
op: *op,
right: Expr::Value(value.clone()),
});
}
PublicTermSql::Null { column, negated } => {
ident(column)?;
preds.push(Predicate::Null {
expr: Self::col_expr(column, qualifier),
negated: *negated,
});
}
}
}
Ok(match preds.len() {
0 => None,
1 => Some(preds.pop().unwrap()),
_ => Some(Predicate::And(preds)),
})
}
fn col_expr(column: &str, qualifier: Option<&str>) -> Expr {
Expr::Column(match qualifier {
Some(q) => format!("{q}.{column}"),
None => column.to_string(),
})
}
fn tenant_pred(
&self,
column: &str,
qualifier: Option<&str>,
) -> Result<Option<Predicate>, OrmError> {
let is_null = Predicate::Null {
expr: Self::col_expr(column, qualifier),
negated: false,
};
let eq = |v: SqlValue| Predicate::Cmp {
left: Self::col_expr(column, qualifier),
op: CmpOp::Eq,
right: Expr::Value(v),
};
Ok(match self.mode {
ScopeMode::All => None,
ScopeMode::NullOnly => Some(is_null),
ScopeMode::Own => {
let v = self.value.clone().ok_or(OrmError::TenancyNoPrincipal)?;
Some(eq(v))
}
ScopeMode::OwnOrNull => {
let v = self.value.clone().ok_or(OrmError::TenancyNoPrincipal)?;
Some(Predicate::Or(vec![eq(v), is_null]))
}
})
}
fn disjunct_pred(
&self,
tenant_col: &str,
session_col: &str,
qualifier: Option<&str>,
) -> Result<Option<Predicate>, OrmError> {
if matches!(self.mode, ScopeMode::All) {
return Ok(None);
}
let eq = |column: &str, v: SqlValue| Predicate::Cmp {
left: Self::col_expr(column, qualifier),
op: CmpOp::Eq,
right: Expr::Value(v),
};
let mut arms = Vec::new();
if let Some(v) = self.value.clone() {
arms.push(eq(tenant_col, v));
}
if let Some(s) = self.session.clone() {
arms.push(eq(session_col, s));
}
match arms.len() {
0 => Err(OrmError::TenancyNoPrincipal),
1 => Ok(arms.pop()),
_ => Ok(Some(Predicate::Or(arms))),
}
}
fn read_pred(
&self,
table: &str,
qualifier: Option<&str>,
) -> Result<Option<Predicate>, OrmError> {
let tenant = match self.resolve_table(table)? {
ResolvedScope::Column(col) => {
ident(&col)?;
self.tenant_pred(&col, qualifier)?
}
ResolvedScope::Unscoped => None,
ResolvedScope::TenantOrSession { tenant, session } => {
ident(&tenant)?;
ident(&session)?;
self.disjunct_pred(&tenant, &session, qualifier)?
}
};
let public = self.public_pred(table, qualifier)?;
let mut out = Predicate::And(Vec::new());
conjoin_front(&mut out, tenant);
conjoin_front(&mut out, public);
Ok(match out {
Predicate::And(v) if v.is_empty() => None,
p => Some(p),
})
}
fn write_target(&self, table: &str) -> Result<Option<(String, SqlValue)>, OrmError> {
let tenant_stamp = || -> Result<Option<SqlValue>, OrmError> {
Ok(match self.mode {
ScopeMode::All => None,
ScopeMode::NullOnly => Some(SqlValue::Null),
ScopeMode::Own | ScopeMode::OwnOrNull => {
Some(self.value.clone().ok_or(OrmError::TenancyNoPrincipal)?)
}
})
};
match self.resolve_table(table)? {
ResolvedScope::Column(col) => Ok(tenant_stamp()?.map(|v| (col, v))),
ResolvedScope::Unscoped => Err(OrmError::UnscopedWrite(table.to_string())),
ResolvedScope::TenantOrSession { tenant, session } => {
if matches!(self.mode, ScopeMode::All) {
return Ok(None);
}
if let Some(v) = self.value.clone() {
Ok(Some((tenant, v)))
} else if let Some(s) = self.session.clone() {
Ok(Some((session, s)))
} else {
Err(OrmError::TenancyNoPrincipal)
}
}
}
}
pub fn is_target(&self) -> bool {
matches!(self.keys, TableKeys::PerTableTarget { .. })
}
fn target_write_allowlist(&self) -> Option<&std::collections::BTreeSet<String>> {
match &self.keys {
TableKeys::PerTableTarget { write, .. } => Some(write),
_ => None,
}
}
fn public_force_cells(&self, table: &str) -> Result<Vec<(String, SqlValue)>, OrmError> {
let TableKeys::PerTableTarget {
public,
require_public,
..
} = &self.keys
else {
return Ok(Vec::new());
};
let terms = match public.get(table) {
Some(t) => t,
None if !require_public => return Ok(Vec::new()),
None => return Err(OrmError::PublicSubsetUndeclared(table.to_string())),
};
let mut out = Vec::with_capacity(terms.len());
for term in terms {
match term {
PublicTermSql::Cmp {
column,
op: CmpOp::Eq,
value,
} => {
ident(column)?;
out.push((column.clone(), value.clone()));
}
PublicTermSql::Null {
column,
negated: false,
} => {
ident(column)?;
out.push((column.clone(), SqlValue::Null));
}
PublicTermSql::Cmp { .. } | PublicTermSql::Null { .. } => {
return Err(OrmError::PublicSubsetNotForceable(table.to_string()))
}
}
}
Ok(out)
}
fn assert_target_settable(&self, table: &str, column: &str) -> Result<(), OrmError> {
let TableKeys::PerTableTarget {
keys,
public,
write,
..
} = &self.keys
else {
return Ok(());
};
let denied = || OrmError::TargetWriteColumnDenied(column.to_string());
if column.contains('.') {
return Err(denied());
}
if !write.iter().any(|c| same_col(c, column)) {
return Err(denied());
}
if let Some(rs) = keys.get(table) {
let tenant_cols: &[&str] = match rs {
ResolvedScope::Column(c) => &[c],
ResolvedScope::TenantOrSession { tenant, session } => &[tenant, session],
ResolvedScope::Unscoped => &[],
};
if tenant_cols.iter().any(|t| same_col(t, column)) {
return Err(denied());
}
}
if let Some(terms) = public.get(table) {
let is_public_col = terms.iter().any(|t| match t {
PublicTermSql::Cmp { column: c, .. } | PublicTermSql::Null { column: c, .. } => {
same_col(c, column)
}
});
if is_public_col {
return Err(denied());
}
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Select {
pub table: String,
pub table_alias: Option<String>,
pub columns: Vec<SelectItem>,
pub joins: Vec<Join>,
pub filter: Option<Predicate>,
pub scope: Option<Scope>,
pub group_by: Vec<Expr>,
pub having: Option<Predicate>,
pub distinct: bool,
pub distinct_on: Vec<Expr>,
pub order: Vec<OrderBy>,
pub limit: Option<u32>,
pub offset: Option<u32>,
pub union: Option<Box<Union>>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Union {
pub all: bool,
pub query: Select,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Assignment {
pub column: String,
pub value: Expr,
}
#[derive(Debug, Clone, PartialEq)]
pub struct RowValues {
pub cells: Vec<Assignment>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct OnConflict {
pub conflict_columns: Vec<String>,
pub update: Vec<Assignment>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Insert {
pub table: String,
pub rows: Vec<RowValues>,
pub conflict: Option<OnConflict>,
pub scope: Option<Scope>,
pub returning: Vec<SelectItem>,
pub from_select: Option<(Vec<String>, Box<Select>)>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Update {
pub table: String,
pub set: Vec<Assignment>,
pub filter: Predicate,
pub scope: Option<Scope>,
pub returning: Vec<SelectItem>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Delete {
pub table: String,
pub filter: Predicate,
pub scope: Option<Scope>,
pub returning: Vec<SelectItem>,
}
impl Select {
pub fn force_scope(&mut self, scope: &Scope) -> Result<(), OrmError> {
self.scope = Some(scope.clone());
self.inject_subquery_scope(scope)?;
if let Some(u) = self.union.as_mut() {
u.query.force_scope(scope)?;
}
Ok(())
}
}
impl Insert {
pub fn force_scope(
&mut self,
write: Option<&Scope>,
read: Option<&Scope>,
) -> Result<(), OrmError> {
if let Some(w) = write {
if let Some(allow) = w.target_write_allowlist() {
self.confine_target_insert(w, allow.is_empty())?;
}
}
self.scope = write.cloned();
let target: Option<(String, SqlValue)> = match write {
Some(w) => w.write_target(&self.table)?,
None => None,
};
if let Some(r) = read {
for row in &mut self.rows {
for cell in &mut row.cells {
inject_scope_expr(r, &mut cell.value)?;
}
}
if let Some(c) = self.conflict.as_mut() {
for a in &mut c.update {
inject_scope_expr(r, &mut a.value)?;
}
}
for it in &mut self.returning {
inject_scope_expr(r, &mut it.expr)?;
}
}
if let Some((cols, src)) = self.from_select.as_mut() {
match read {
Some(r) => src.force_scope(r)?,
None => src.scope = None,
}
if let Some((column, v)) = &target {
let column = column.clone();
if let Some(i) = cols.iter().position(|c| same_col(c, &column)) {
cols.remove(i);
drop_projection_at(src, i);
}
let v = v.clone();
cols.push(column);
push_projection(
src,
SelectItem {
expr: Expr::Value(v),
alias: None,
},
);
}
}
Ok(())
}
fn confine_target_insert(
&mut self,
scope: &Scope,
empty_allowlist: bool,
) -> Result<(), OrmError> {
if empty_allowlist {
return Err(OrmError::TargetWriteNotGranted(self.table.clone()));
}
if self.from_select.is_some() {
return Err(OrmError::TargetWriteUnsupported("INSERT … SELECT"));
}
if self.conflict.is_some() {
return Err(OrmError::TargetWriteUnsupported("ON CONFLICT upsert"));
}
for row in &self.rows {
for cell in &row.cells {
scope.assert_target_settable(&self.table, &cell.column)?;
}
}
let forced = scope.public_force_cells(&self.table)?;
for row in &mut self.rows {
for (column, value) in &forced {
row.cells.push(Assignment {
column: column.clone(),
value: Expr::Value(value.clone()),
});
}
}
Ok(())
}
}
fn drop_projection_at(s: &mut Select, i: usize) {
if i < s.columns.len() {
s.columns.remove(i);
}
if let Some(u) = s.union.as_mut() {
drop_projection_at(&mut u.query, i);
}
}
fn push_projection(s: &mut Select, item: SelectItem) {
s.columns.push(item.clone());
if let Some(u) = s.union.as_mut() {
push_projection(&mut u.query, item);
}
}
fn conjoin_front(filter: &mut Predicate, add: Option<Predicate>) {
let Some(a) = add else { return };
if matches!(filter, Predicate::And(v) if v.is_empty()) {
*filter = a;
} else {
let existing = std::mem::replace(filter, Predicate::And(Vec::new()));
*filter = Predicate::And(vec![a, existing]);
}
}
fn own_rank_expr(scope: &Scope) -> Expr {
let Some(value) = scope.value.clone() else {
return Expr::Value(SqlValue::Integer(0));
};
let col = || Expr::Column(scope.column.clone());
let own = Predicate::And(vec![
Predicate::Null {
expr: col(),
negated: true,
},
Predicate::Cmp {
left: col(),
op: CmpOp::Eq,
right: Expr::Value(value),
},
]);
Expr::Case {
branches: vec![(own, Expr::Value(SqlValue::Integer(1)))],
otherwise: Some(Box::new(Expr::Value(SqlValue::Integer(0)))),
}
}
fn conjoin_subquery_scope(
scope: &Scope,
table: &str,
filter: &mut Predicate,
) -> Result<(), OrmError> {
conjoin_front(filter, scope.read_pred(table, Some(table))?);
Ok(())
}
fn inject_scope_expr(scope: &Scope, e: &mut Expr) -> Result<(), OrmError> {
match e {
Expr::IsOwn => *e = own_rank_expr(scope),
Expr::RelatedAggregate { table, filter, .. }
| Expr::RelatedScalar { table, filter, .. } => {
inject_scope_pred(scope, filter)?;
conjoin_subquery_scope(scope, table, filter)?;
}
Expr::Aggregate(_, inner) | Expr::JsonExtract(inner, _) => inject_scope_expr(scope, inner)?,
Expr::Binary(_, l, r) | Expr::JsonExtractDyn(l, r) | Expr::JsonConcat(l, r) => {
inject_scope_expr(scope, l)?;
inject_scope_expr(scope, r)?;
}
Expr::Distance { left, right, .. } => {
inject_scope_expr(scope, left)?;
inject_scope_expr(scope, right)?;
}
Expr::Func(_, args) => {
for a in args.iter_mut() {
inject_scope_expr(scope, a)?;
}
}
Expr::Case {
branches,
otherwise,
} => {
for (when, then) in branches {
inject_scope_pred(scope, when)?;
inject_scope_expr(scope, then)?;
}
if let Some(e) = otherwise {
inject_scope_expr(scope, e)?;
}
}
Expr::Column(_) | Expr::Value(_) | Expr::Star | Expr::VectorLiteral(_) => {}
}
Ok(())
}
fn inject_scope_pred(scope: &Scope, p: &mut Predicate) -> Result<(), OrmError> {
match p {
Predicate::InSubquery {
expr,
table,
filter,
..
} => {
inject_scope_expr(scope, expr)?;
inject_scope_pred(scope, filter)?;
conjoin_subquery_scope(scope, table, filter)?;
}
Predicate::And(v) | Predicate::Or(v) => {
for c in v.iter_mut() {
inject_scope_pred(scope, c)?;
}
}
Predicate::Not(inner) => inject_scope_pred(scope, inner)?,
Predicate::Cmp { left, right, .. } => {
inject_scope_expr(scope, left)?;
inject_scope_expr(scope, right)?;
}
Predicate::Between {
expr, low, high, ..
} => {
inject_scope_expr(scope, expr)?;
inject_scope_expr(scope, low)?;
inject_scope_expr(scope, high)?;
}
Predicate::In { expr, values, .. } => {
inject_scope_expr(scope, expr)?;
for v in values.iter_mut() {
inject_scope_expr(scope, v)?;
}
}
Predicate::Like { expr, .. } | Predicate::Null { expr, .. } => {
inject_scope_expr(scope, expr)?;
}
}
Ok(())
}
impl Select {
fn inject_subquery_scope(&mut self, scope: &Scope) -> Result<(), OrmError> {
for it in &mut self.columns {
inject_scope_expr(scope, &mut it.expr)?;
}
for e in &mut self.distinct_on {
inject_scope_expr(scope, e)?;
}
if let Some(f) = self.filter.as_mut() {
inject_scope_pred(scope, f)?;
}
if let Some(h) = self.having.as_mut() {
inject_scope_pred(scope, h)?;
}
for e in &mut self.group_by {
inject_scope_expr(scope, e)?;
}
for o in &mut self.order {
inject_scope_expr(scope, &mut o.expr)?;
}
for j in &mut self.joins {
inject_scope_pred(scope, &mut j.on)?;
}
Ok(())
}
}
impl Update {
pub fn force_scope(&mut self, scope: &Scope) -> Result<(), OrmError> {
if let Some(allow) = scope.target_write_allowlist() {
if allow.is_empty() {
return Err(OrmError::TargetWriteNotGranted(self.table.clone()));
}
for a in &self.set {
scope.assert_target_settable(&self.table, &a.column)?;
}
if let Some(pred) = scope.public_pred(&self.table, None)? {
conjoin_front(&mut self.filter, Some(pred));
}
}
self.scope = Some(scope.clone());
for a in &mut self.set {
inject_scope_expr(scope, &mut a.value)?;
}
inject_scope_pred(scope, &mut self.filter)?;
for it in &mut self.returning {
inject_scope_expr(scope, &mut it.expr)?;
}
Ok(())
}
}
impl Delete {
pub fn force_scope(&mut self, scope: &Scope) -> Result<(), OrmError> {
if scope.is_target() {
return Err(OrmError::TargetDeleteRefused(self.table.clone()));
}
self.scope = Some(scope.clone());
inject_scope_pred(scope, &mut self.filter)?;
for it in &mut self.returning {
inject_scope_expr(scope, &mut it.expr)?;
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, thiserror::Error)]
pub enum OrmError {
#[error("invalid identifier: {0:?}")]
InvalidIdentifier(String),
#[error("empty query: {0}")]
Empty(&'static str),
#[error("bad expression: {0}")]
BadExpr(&'static str),
#[error("tenancy: table {0:?} has no declared scope (deny-by-default)")]
TenancyUndeclared(String),
#[error("tenancy: table {0:?} is Unscoped (global reference); guest writes are refused (deny-by-default)")]
UnscopedWrite(String),
#[error("tenancy: no resolved principal for a scoped operation (deny-by-default)")]
TenancyNoPrincipal,
#[error(
"tenancy: table {0:?} has no declared public subset for a target read (deny-by-default)"
)]
PublicSubsetUndeclared(String),
#[error("tenancy: target route {0:?} has no write grant (read-only; deny-by-default)")]
TargetWriteNotGranted(String),
#[error("tenancy: target write may not set column {0:?} (not in the write allowlist)")]
TargetWriteColumnDenied(String),
#[error("tenancy: a target-tenant DELETE is refused (target writes are INSERT/UPDATE only)")]
TargetDeleteRefused(String),
#[error("tenancy: target INSERT cannot force table {0:?} into its public subset (a non-equality/non-null public term); refused")]
PublicSubsetNotForceable(String),
#[error("tenancy: unsupported target write shape ({0}); target writes are a plain INSERT or a confined UPDATE only")]
TargetWriteUnsupported(&'static str),
}
pub type Compiled = (String, Vec<SqlValue>);
fn ident(name: &str) -> Result<&str, OrmError> {
let ok = |s: &str| {
let mut cs = s.chars();
matches!(cs.next(), Some(c) if c == '_' || c.is_ascii_alphabetic())
&& s.chars().all(|c| c == '_' || c.is_ascii_alphanumeric())
};
let valid = match name.split_once('.') {
Some((t, c)) => !t.is_empty() && !c.is_empty() && ok(t) && ok(c),
None => ok(name),
};
if valid {
Ok(name)
} else {
Err(OrmError::InvalidIdentifier(name.to_string()))
}
}
fn same_col(a: &str, b: &str) -> bool {
let base = |s: &str| s.rsplit('.').next().unwrap_or(s).to_ascii_lowercase();
base(a) == base(b)
}
#[derive(Default)]
struct Params(Vec<SqlValue>);
impl Params {
fn bind(&mut self, v: SqlValue) -> String {
self.0.push(v);
format!("?{}", self.0.len())
}
}
fn render_expr(e: &Expr, params: &mut Params, dialect: Dialect) -> Result<String, OrmError> {
Ok(match e {
Expr::Column(name) => ident(name)?.to_string(),
Expr::Value(v) => params.bind(v.clone()),
Expr::Star => {
return Err(OrmError::BadExpr(
"`*` is only valid as the count(*) argument",
))
}
Expr::Aggregate(agg, inner) => {
let arg = match inner.as_ref() {
Expr::Star if *agg == Agg::Count => "*".to_string(),
Expr::Star => return Err(OrmError::BadExpr("`*` is only valid as count(*)")),
other => render_expr(other, params, dialect)?,
};
format!("{}({arg})", agg.keyword())
}
Expr::Binary(op, l, r) => {
format!(
"({} {} {})",
render_expr(l, params, dialect)?,
op.symbol(),
render_expr(r, params, dialect)?
)
}
Expr::Func(f, args) => {
let (name, min, max) = f.spec();
if args.len() < min || max.is_some_and(|m| args.len() > m) {
return Err(OrmError::BadExpr("function called with the wrong arity"));
}
if args.is_empty() {
name.to_string()
} else {
let rendered: Result<Vec<String>, _> = args
.iter()
.map(|a| render_expr(a, params, dialect))
.collect();
format!("{name}({})", rendered?.join(", "))
}
}
Expr::JsonExtract(inner, path) => {
if path.is_empty() {
return Err(OrmError::BadExpr("json extract needs at least one key"));
}
for k in path {
ident(k)?;
}
let base = render_expr(inner, params, dialect)?;
match dialect {
Dialect::Postgres => format!("({base}) #>> '{{{}}}'", path.join(",")),
Dialect::Sqlite | Dialect::Mysql => {
let p = params.bind(SqlValue::Text(format!("$.{}", path.join("."))));
format!("json_extract({base}, {p})")
}
}
}
Expr::Distance {
left,
right,
metric,
} => {
if dialect != Dialect::Postgres {
return Err(OrmError::BadExpr("vector distance is Postgres-only"));
}
format!(
"({} {} {})",
render_expr(left, params, dialect)?,
metric.operator(),
render_expr(right, params, dialect)?,
)
}
Expr::VectorLiteral(v) => {
if dialect != Dialect::Postgres {
return Err(OrmError::BadExpr("vector literals are Postgres-only"));
}
let p = params.bind(SqlValue::Text(vector_literal(v)?));
format!("{p}::vector")
}
Expr::RelatedAggregate {
agg,
arg,
table,
filter,
} => {
let arg_sql = match arg {
RelArg::Star if *agg == Agg::Count => "*".to_string(),
RelArg::Star => return Err(OrmError::BadExpr("`*` is only valid as count(*)")),
RelArg::Column(c) => ident(c)?.to_string(),
};
let table_sql = ident(table)?;
let where_sql = render_pred(filter, params, false, dialect)?;
format!(
"(SELECT {}({arg_sql}) FROM {table_sql} WHERE {where_sql})",
agg.keyword()
)
}
Expr::RelatedScalar {
column,
table,
filter,
} => {
let col_sql = ident(column)?;
let table_sql = ident(table)?;
let where_sql = render_pred(filter, params, false, dialect)?;
format!("(SELECT {col_sql} FROM {table_sql} WHERE {where_sql})")
}
Expr::JsonExtractDyn(base, key) => {
if matches!(dialect, Dialect::Mysql) {
return Err(OrmError::BadExpr(
"dynamic-key json extract (->> <bound>) is not supported on MySQL",
));
}
format!(
"({} ->> {})",
render_expr(base, params, dialect)?,
render_expr(key, params, dialect)?,
)
}
Expr::JsonConcat(left, right) => {
if dialect != Dialect::Postgres {
return Err(OrmError::BadExpr("json concat (||) is Postgres-only"));
}
format!(
"({} || {})",
render_expr(left, params, dialect)?,
render_expr(right, params, dialect)?,
)
}
Expr::Case {
branches,
otherwise,
} => {
if branches.is_empty() {
return Err(OrmError::BadExpr("CASE has no WHEN branches"));
}
let mut s = String::from("CASE");
for (when, then) in branches {
let w = render_pred(when, params, false, dialect)?;
let t = render_expr(then, params, dialect)?;
s.push_str(&format!(" WHEN {w} THEN {t}"));
}
if let Some(e) = otherwise {
let e = render_expr(e, params, dialect)?;
s.push_str(&format!(" ELSE {e}"));
}
s.push_str(" END");
format!("({s})")
}
Expr::IsOwn => {
return Err(OrmError::BadExpr(
"is_own()/own_first() requires an own-tenant (own or own+null) read scope",
))
}
})
}
fn vector_literal(s: &str) -> Result<String, OrmError> {
let inner = s
.trim()
.strip_prefix('[')
.and_then(|x| x.strip_suffix(']'))
.ok_or(OrmError::BadExpr(
"vector literal must be a bracketed list like [0.1, 0.2]",
))?;
if inner.trim().is_empty() {
return Err(OrmError::BadExpr(
"vector literal must have at least one component",
));
}
let mut parts = Vec::new();
for part in inner.split(',') {
let p = part.trim();
let f: f64 = p
.parse()
.map_err(|_| OrmError::BadExpr("vector literal component is not a number"))?;
if !f.is_finite() {
return Err(OrmError::BadExpr("vector literal component must be finite"));
}
parts.push(p);
}
Ok(format!("[{}]", parts.join(",")))
}
fn render_pred(
p: &Predicate,
params: &mut Params,
nested: bool,
dialect: Dialect,
) -> Result<String, OrmError> {
let compound = |body: String| {
if nested {
format!("({body})")
} else {
body
}
};
Ok(match p {
Predicate::And(ps) => {
if ps.is_empty() {
"1 = 1".to_string()
} else {
let parts: Result<Vec<String>, _> = ps
.iter()
.map(|c| render_pred(c, params, true, dialect))
.collect();
compound(parts?.join(" AND "))
}
}
Predicate::Or(ps) => {
if ps.is_empty() {
"1 = 0".to_string()
} else {
let parts: Result<Vec<String>, _> = ps
.iter()
.map(|c| render_pred(c, params, true, dialect))
.collect();
compound(parts?.join(" OR "))
}
}
Predicate::Not(inner) => format!("NOT {}", render_pred(inner, params, true, dialect)?),
Predicate::Cmp { left, op, right } => format!(
"{} {} {}",
render_expr(left, params, dialect)?,
op.symbol(),
render_expr(right, params, dialect)?
),
Predicate::Between {
expr,
low,
high,
negated,
} => format!(
"{} {}BETWEEN {} AND {}",
render_expr(expr, params, dialect)?,
if *negated { "NOT " } else { "" },
render_expr(low, params, dialect)?,
render_expr(high, params, dialect)?
),
Predicate::In {
expr,
values,
negated,
} => {
if values.is_empty() {
if *negated { "1 = 1" } else { "1 = 0" }.to_string()
} else {
let lhs = render_expr(expr, params, dialect)?;
let ph: Result<Vec<String>, _> = values
.iter()
.map(|v| render_expr(v, params, dialect))
.collect();
format!(
"{lhs} {}IN ({})",
if *negated { "NOT " } else { "" },
ph?.join(", ")
)
}
}
Predicate::Like {
expr,
pattern,
insensitive,
negated,
} => {
let neg = if *negated { "NOT " } else { "" };
let lhs = render_expr(expr, params, dialect)?;
let pat = params.bind(SqlValue::Text(pattern.clone()));
if *insensitive {
format!("lower({lhs}) {neg}LIKE lower({pat})")
} else {
format!("{lhs} {neg}LIKE {pat}")
}
}
Predicate::Null { expr, negated } => format!(
"{} IS {}NULL",
render_expr(expr, params, dialect)?,
if *negated { "NOT " } else { "" }
),
Predicate::InSubquery {
expr,
column,
table,
filter,
negated,
} => {
let lhs = render_expr(expr, params, dialect)?;
let col_sql = ident(column)?;
let table_sql = ident(table)?;
let where_sql = render_pred(filter, params, false, dialect)?;
let not = if *negated { "NOT " } else { "" };
format!("{lhs} {not}IN (SELECT {col_sql} FROM {table_sql} WHERE {where_sql})")
}
})
}
fn render_where(
scope_pred: Option<Predicate>,
filter: Option<&Predicate>,
params: &mut Params,
dialect: Dialect,
) -> Result<Option<String>, OrmError> {
let filter = filter.filter(|f| !matches!(f, Predicate::And(v) if v.is_empty()));
let combined = match (scope_pred, filter) {
(None, None) => return Ok(None),
(Some(s), None) => s,
(None, Some(f)) => f.clone(),
(Some(s), Some(f)) => Predicate::And(vec![s, f.clone()]),
};
Ok(Some(render_pred(&combined, params, false, dialect)?))
}
fn single_scope_pred(scope: Option<&Scope>, table: &str) -> Result<Option<Predicate>, OrmError> {
let Some(s) = scope else { return Ok(None) };
match s.write_target(table)? {
None => Ok(None), Some((col, value)) => {
ident(&col)?;
let col_expr = Expr::Column(col);
let pred = if matches!(value, SqlValue::Null) {
Predicate::Null {
expr: col_expr,
negated: false,
}
} else {
Predicate::Cmp {
left: col_expr,
op: CmpOp::Eq,
right: Expr::Value(value),
}
};
Ok(Some(pred))
}
}
}
fn render_select_items(
items: &[SelectItem],
params: &mut Params,
dialect: Dialect,
) -> Result<String, OrmError> {
if items.is_empty() {
return Ok("*".to_string());
}
let parts: Result<Vec<String>, _> = items
.iter()
.map(|it| {
let e = render_expr(&it.expr, params, dialect)?;
Ok::<String, OrmError>(match &it.alias {
Some(a) => format!("{e} AS {}", ident(a)?),
None => e,
})
})
.collect();
Ok(parts?.join(", "))
}
fn render_returning(
items: &[SelectItem],
params: &mut Params,
dialect: Dialect,
) -> Result<String, OrmError> {
if items.is_empty() {
Ok(String::new())
} else {
Ok(format!(
" RETURNING {}",
render_select_items(items, params, dialect)?
))
}
}
impl Select {
pub fn from(table: impl Into<String>) -> Self {
Self {
table: table.into(),
table_alias: None,
columns: Vec::new(),
joins: Vec::new(),
filter: None,
scope: None,
group_by: Vec::new(),
having: None,
distinct: false,
distinct_on: Vec::new(),
order: Vec::new(),
limit: None,
offset: None,
union: None,
}
}
pub fn compile(&self, dialect: Dialect) -> Result<Compiled, OrmError> {
let mut params = Params::default();
let sql = self.render_into(&mut params, dialect)?;
Ok((sql, params.0))
}
fn scope_where_pred(&self) -> Result<Option<Predicate>, OrmError> {
let Some(scope) = &self.scope else {
return Ok(None);
};
if self.joins.is_empty() {
return scope.read_pred(&self.table, None);
}
let refs: Vec<(&str, &str)> = std::iter::once((
self.table.as_str(),
self.table_alias.as_deref().unwrap_or(&self.table),
))
.chain(
self.joins
.iter()
.map(|j| (j.table.as_str(), j.alias.as_deref().unwrap_or(&j.table))),
)
.collect();
let mut parts: Vec<Predicate> = Vec::with_capacity(refs.len());
for (table, qual) in refs {
ident(qual)?;
if let Some(p) = scope.read_pred(table, Some(qual))? {
parts.push(p);
}
}
Ok((!parts.is_empty()).then_some(Predicate::And(parts)))
}
fn render_into(&self, params: &mut Params, dialect: Dialect) -> Result<String, OrmError> {
let mut sql = self.render_body(params, dialect)?;
if let Some(u) = &self.union {
let kw = if u.all { "UNION ALL" } else { "UNION" };
let branch = u.query.render_body(params, dialect)?;
sql.push_str(&format!(" {kw} {branch}"));
}
Ok(sql)
}
fn render_body(&self, params: &mut Params, dialect: Dialect) -> Result<String, OrmError> {
let table = ident(&self.table)?;
let distinct = if !self.distinct_on.is_empty() {
if dialect != Dialect::Postgres {
return Err(OrmError::BadExpr("DISTINCT ON is Postgres-only"));
}
let cols = self
.distinct_on
.iter()
.map(|e| render_expr(e, &mut *params, dialect))
.collect::<Result<Vec<_>, _>>()?;
format!("DISTINCT ON ({}) ", cols.join(", "))
} else if self.distinct {
"DISTINCT ".to_string()
} else {
String::new()
};
let select_list = render_select_items(&self.columns, &mut *params, dialect)?;
let mut sql = format!("SELECT {distinct}{select_list} FROM {table}");
if let Some(a) = &self.table_alias {
sql.push_str(&format!(" AS {}", ident(a)?));
}
for j in &self.joins {
let jt = ident(&j.table)?;
let kw = match j.kind {
JoinKind::Inner => "JOIN",
JoinKind::Left => "LEFT JOIN",
};
sql.push_str(&format!(" {kw} {jt}"));
if let Some(a) = &j.alias {
sql.push_str(&format!(" AS {}", ident(a)?));
}
sql.push_str(&format!(
" ON {}",
render_pred(&j.on, &mut *params, false, dialect)?
));
}
if let Some(w) = render_where(
self.scope_where_pred()?,
self.filter.as_ref(),
&mut *params,
dialect,
)? {
sql.push_str(&format!(" WHERE {w}"));
}
if !self.group_by.is_empty() {
let terms: Result<Vec<String>, _> = self
.group_by
.iter()
.map(|e| render_expr(e, &mut *params, dialect))
.collect();
sql.push_str(&format!(" GROUP BY {}", terms?.join(", ")));
}
if let Some(h) = &self.having {
sql.push_str(&format!(
" HAVING {}",
render_pred(h, &mut *params, false, dialect)?
));
}
if !self.order.is_empty() {
let terms: Result<Vec<String>, _> = self
.order
.iter()
.map(|o| {
let e = render_expr(&o.expr, &mut *params, dialect)?;
let d = match o.dir {
Direction::Asc => "ASC",
Direction::Desc => "DESC",
};
Ok::<String, OrmError>(format!("{e} {d}"))
})
.collect();
sql.push_str(&format!(" ORDER BY {}", terms?.join(", ")));
}
if let Some(n) = self.limit {
sql.push_str(&format!(" LIMIT {n}"));
}
if let Some(n) = self.offset {
sql.push_str(&format!(" OFFSET {n}"));
}
Ok(sql)
}
}
impl Insert {
pub fn compile(&self, dialect: Dialect) -> Result<Compiled, OrmError> {
let table = ident(&self.table)?;
let mut params = Params::default();
let stamp: Option<(String, SqlValue)> = match self.scope.as_ref() {
Some(s) => s.write_target(&self.table)?,
None => None,
};
if let Some((cols, select)) = &self.from_select {
let col_sql = cols
.iter()
.map(|c| ident(c).map(str::to_string))
.collect::<Result<Vec<_>, _>>()?;
if col_sql.is_empty() {
return Err(OrmError::Empty("insert-select has no columns"));
}
let select_sql = select.render_into(&mut params, dialect)?;
let mut sql = format!("INSERT INTO {table} ({}) {select_sql}", col_sql.join(", "));
sql.push_str(&render_conflict(
self.conflict.as_ref(),
stamp.as_ref().map(|(c, v)| (c.as_str(), v)),
&mut params,
dialect,
)?);
sql.push_str(&render_returning(&self.returning, &mut params, dialect)?);
return Ok((sql, params.0));
}
if self.rows.is_empty() {
return Err(OrmError::Empty("insert has no rows"));
}
let mut columns: Vec<String> = Vec::new();
for a in &self.rows[0].cells {
let c = ident(&a.column)?.to_string();
if !columns.contains(&c) {
columns.push(c);
}
}
if let Some((column, _)) = &stamp {
let c = ident(column)?.to_string();
if !columns.iter().any(|existing| same_col(existing, &c)) {
columns.push(c);
}
}
if columns.is_empty() {
return Err(OrmError::Empty("insert row has no columns"));
}
let mut value_groups: Vec<String> = Vec::new();
for row in &self.rows {
let mut ph: Vec<String> = Vec::with_capacity(columns.len());
for col in &columns {
if let Some((column, value)) = &stamp {
if same_col(column, col) {
ph.push(params.bind(value.clone()));
continue;
}
}
match row.cells.iter().find(|a| same_col(&a.column, col)) {
Some(a) => ph.push(render_expr(&a.value, &mut params, dialect)?),
None => ph.push(params.bind(SqlValue::Null)),
}
}
value_groups.push(format!("({})", ph.join(", ")));
}
let mut sql = format!(
"INSERT INTO {table} ({}) VALUES {}",
columns.join(", "),
value_groups.join(", ")
);
sql.push_str(&render_conflict(
self.conflict.as_ref(),
stamp.as_ref().map(|(c, v)| (c.as_str(), v)),
&mut params,
dialect,
)?);
sql.push_str(&render_returning(&self.returning, &mut params, dialect)?);
Ok((sql, params.0))
}
}
fn render_conflict(
conflict: Option<&OnConflict>,
stamp: Option<(&str, &SqlValue)>,
params: &mut Params,
dialect: Dialect,
) -> Result<String, OrmError> {
let Some(oc) = conflict else {
return Ok(String::new());
};
let conflict_cols = oc
.conflict_columns
.iter()
.map(|c| ident(c).map(str::to_string))
.collect::<Result<Vec<_>, _>>()?;
let guard = stamp;
let do_nothing = || format!(" ON CONFLICT ({}) DO NOTHING", conflict_cols.join(", "));
if oc.update.is_empty() {
return Ok(do_nothing());
}
if guard.is_some() && matches!(dialect, Dialect::Mysql) {
return Err(OrmError::BadExpr(
"a tenant-scoped upsert (ON CONFLICT DO UPDATE) is unsupported on MySQL \
(ON DUPLICATE KEY UPDATE cannot be bounded to the tenant's rows)",
));
}
let sets = oc
.update
.iter()
.filter(|a| guard.is_none_or(|(col, _)| !same_col(&a.column, col)))
.map(|a| {
let c = ident(&a.column)?;
Ok::<String, OrmError>(format!("{c} = {}", render_expr(&a.value, params, dialect)?))
})
.collect::<Result<Vec<_>, _>>()?;
if sets.is_empty() {
return Ok(do_nothing());
}
let mut clause = format!(
" ON CONFLICT ({}) DO UPDATE SET {}",
conflict_cols.join(", "),
sets.join(", ")
);
if let Some((col, value)) = guard {
ident(col)?;
let col_expr = Expr::Column(col.to_string());
let pred = if matches!(value, SqlValue::Null) {
Predicate::Null {
expr: col_expr,
negated: false,
}
} else {
Predicate::Cmp {
left: col_expr,
op: CmpOp::Eq,
right: Expr::Value(value.clone()),
}
};
clause.push_str(&format!(
" WHERE {}",
render_pred(&pred, params, false, dialect)?
));
}
Ok(clause)
}
impl Update {
pub fn compile(&self, dialect: Dialect) -> Result<Compiled, OrmError> {
if self.set.is_empty() {
return Err(OrmError::Empty("update has no assignments"));
}
let empty_filter =
matches!(&self.filter, Predicate::And(v) | Predicate::Or(v) if v.is_empty());
if empty_filter && self.scope.is_none() {
return Err(OrmError::Empty(
"update has an empty filter (unbounded update refused)",
));
}
let table = ident(&self.table)?;
let mut params = Params::default();
let scope_col: Option<String> = match self.scope.as_ref() {
Some(s) => s.write_target(&self.table)?.map(|(col, _)| col),
None => None,
};
let sets: Result<Vec<String>, _> = self
.set
.iter()
.filter(|a| {
scope_col
.as_deref()
.is_none_or(|col| !same_col(&a.column, col))
})
.map(|a| {
let c = ident(&a.column)?;
Ok::<String, OrmError>(format!(
"{c} = {}",
render_expr(&a.value, &mut params, dialect)?
))
})
.collect();
let sets = sets?;
if sets.is_empty() {
return Err(OrmError::Empty(
"update has no assignments left after dropping the tenant column",
));
}
let set_sql = sets.join(", ");
let where_sql = render_where(
single_scope_pred(self.scope.as_ref(), &self.table)?,
Some(&self.filter),
&mut params,
dialect,
)?
.ok_or(OrmError::Empty(
"update has an empty filter (unbounded update refused)",
))?;
let mut sql = format!("UPDATE {table} SET {set_sql} WHERE {where_sql}");
sql.push_str(&render_returning(&self.returning, &mut params, dialect)?);
Ok((sql, params.0))
}
}
impl Delete {
pub fn compile(&self, dialect: Dialect) -> Result<Compiled, OrmError> {
let empty_filter =
matches!(&self.filter, Predicate::And(v) | Predicate::Or(v) if v.is_empty());
if empty_filter && self.scope.is_none() {
return Err(OrmError::Empty(
"delete has an empty filter (unbounded delete refused)",
));
}
let table = ident(&self.table)?;
let mut params = Params::default();
let where_sql = render_where(
single_scope_pred(self.scope.as_ref(), &self.table)?,
Some(&self.filter),
&mut params,
dialect,
)?
.ok_or(OrmError::Empty(
"delete has an empty filter (unbounded delete refused)",
))?;
let mut sql = format!("DELETE FROM {table} WHERE {where_sql}");
sql.push_str(&render_returning(&self.returning, &mut params, dialect)?);
Ok((sql, params.0))
}
}
pub fn compile_promote(scope: &Scope, table: &str, dialect: Dialect) -> Result<Compiled, OrmError> {
if scope.is_target() {
return Err(OrmError::TargetWriteUnsupported("promote"));
}
let (tenant_col, session_col) = match scope.resolve_table(table)? {
ResolvedScope::TenantOrSession { tenant, session } => (tenant, session),
_ => {
return Err(OrmError::BadExpr(
"promote requires a TenantOrSession table (an anonymous-first table)",
))
}
};
ident(&tenant_col)?;
ident(&session_col)?;
let tenant = scope.value.clone().ok_or(OrmError::TenancyNoPrincipal)?;
let session = scope.session.clone().ok_or(OrmError::TenancyNoPrincipal)?;
let promote = Update {
table: table.to_string(),
set: vec![Assignment {
column: tenant_col.clone(),
value: Expr::val(tenant),
}],
filter: Predicate::And(vec![
Predicate::Cmp {
left: Expr::Column(session_col),
op: CmpOp::Eq,
right: Expr::val(session),
},
Predicate::Null {
expr: Expr::Column(tenant_col),
negated: false,
},
]),
scope: None,
returning: vec![],
};
promote.compile(dialect)
}
#[derive(Debug, Clone, PartialEq)]
pub struct AttachReference {
pub child: String,
pub parent: String,
pub ref_column: String,
pub ref_value: SqlValue,
pub set: Vec<Assignment>,
}
pub fn compile_attach_reference(
scope: &Scope,
spec: &AttachReference,
dialect: Dialect,
) -> Result<Compiled, OrmError> {
ident(&spec.ref_column)?;
let child_tenant = match scope.resolve_table(&spec.child)? {
ResolvedScope::Column(c) => c,
_ => {
return Err(OrmError::BadExpr(
"attach_reference child must be a plain tenant table",
))
}
};
let parent_tenant = match scope.resolve_table(&spec.parent)? {
ResolvedScope::Column(c) => c,
_ => {
return Err(OrmError::BadExpr(
"attach_reference parent must be a plain tenant table",
))
}
};
ident(&child_tenant)?;
ident(&parent_tenant)?;
let is_target = scope.is_target();
let mut columns: Vec<String> = Vec::with_capacity(spec.set.len() + 2);
let mut projection: Vec<SelectItem> = Vec::with_capacity(spec.set.len() + 2);
for a in &spec.set {
if same_col(&a.column, &child_tenant) {
return Err(OrmError::TargetWriteColumnDenied(a.column.clone()));
}
if is_target {
scope.assert_target_settable(&spec.child, &a.column)?;
} else {
ident(&a.column)?;
}
columns.push(a.column.clone());
projection.push(SelectItem {
expr: a.value.clone(),
alias: None,
});
}
columns.push(child_tenant);
projection.push(SelectItem {
expr: Expr::Column(parent_tenant),
alias: None,
});
if is_target {
for (col, val) in scope.public_force_cells(&spec.child)? {
columns.push(col);
projection.push(SelectItem {
expr: Expr::Value(val),
alias: None,
});
}
}
let mut source = Select {
columns: projection,
filter: Some(Predicate::Cmp {
left: Expr::Column(spec.ref_column.clone()),
op: CmpOp::Eq,
right: Expr::Value(spec.ref_value.clone()),
}),
..Select::from(spec.parent.clone())
};
source.force_scope(scope)?;
let insert = Insert {
table: spec.child.clone(),
rows: vec![],
conflict: None,
scope: None,
returning: vec![],
from_select: Some((columns, Box::new(source))),
};
insert.compile(dialect)
}
#[cfg(test)]
mod tests {
use super::*;
fn t(s: &str) -> SqlValue {
SqlValue::Text(s.to_string())
}
fn cmp(col: &str, op: CmpOp, v: SqlValue) -> Predicate {
Predicate::Cmp {
left: Expr::Column(col.into()),
op,
right: Expr::Value(v),
}
}
fn item(e: Expr) -> SelectItem {
SelectItem {
expr: e,
alias: None,
}
}
#[test]
fn select_basic_where_order_limit() {
let q = Select {
columns: vec![item(Expr::col("id")), item(Expr::col("state"))],
filter: Some(cmp("project_id", CmpOp::Eq, t("prj_1"))),
order: vec![OrderBy {
expr: Expr::col("created_at"),
dir: Direction::Desc,
}],
limit: Some(10),
..Select::from("work_order")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT id, state FROM work_order WHERE project_id = ?1 ORDER BY created_at DESC LIMIT 10"
);
assert_eq!(params, vec![t("prj_1")]);
}
#[test]
fn scope_is_anded_and_bound_first() {
let q = Select {
filter: Some(cmp("kind", CmpOp::Eq, t("supplier"))),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
..Select::from("party")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT * FROM party WHERE tenant_id = ?1 AND kind = ?2"
);
assert_eq!(params, vec![t("ten_1"), t("supplier")]);
}
#[test]
fn per_table_keys_scope_each_ref_on_its_own_column() {
use std::collections::BTreeMap;
let q = Select {
table_alias: Some("sc".into()),
joins: vec![Join {
kind: JoinKind::Left,
table: "tenant".into(),
alias: Some("t".into()),
on: Predicate::Cmp {
left: Expr::col("sc.tenant_id"),
op: CmpOp::Eq,
right: Expr::col("t.id"),
},
}],
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("acme")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([
(
"storefront_config".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
),
(
"tenant".to_string(),
ResolvedScope::Column("id".to_string()),
),
])),
}),
..Select::from("storefront_config")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert!(
sql.contains("sc.tenant_id = ?"),
"base scoped on tenant_id: {sql}"
);
assert!(
sql.contains("t.id = ?"),
"identity table scoped on its own PK: {sql}"
);
assert_eq!(params, vec![t("acme"), t("acme")]);
let mut q2 = Select {
table_alias: Some("sc".into()),
joins: vec![Join {
kind: JoinKind::Left,
table: "countries".into(),
alias: Some("c".into()),
on: Predicate::Cmp {
left: Expr::col("sc.country"),
op: CmpOp::Eq,
right: Expr::col("c.code"),
},
}],
..Select::from("storefront_config")
};
q2.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("acme")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([
(
"storefront_config".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
),
("countries".to_string(), ResolvedScope::Unscoped),
])),
})
.unwrap();
let (sql2, params2) = q2.compile(Dialect::Sqlite).unwrap();
assert!(sql2.contains("sc.tenant_id = ?"), "sql2: {sql2}");
assert_eq!(
params2,
vec![t("acme")],
"unscoped join adds no tenant predicate: {sql2}"
);
let mut q3 = Select::from("secret_table");
q3.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("acme")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([(
"orders".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
)])),
})
.unwrap();
assert!(matches!(
q3.compile(Dialect::Sqlite),
Err(OrmError::TenancyUndeclared(tbl)) if tbl == "secret_table"
));
}
#[test]
fn is_own_lowers_to_a_case_rank_and_orders_own_first() {
let mut q = Select {
columns: vec![item(Expr::col("body"))],
filter: Some(cmp("key_name", CmpOp::Eq, t("k"))),
order: vec![OrderBy {
expr: Expr::IsOwn,
dir: Direction::Desc,
}],
limit: Some(1),
..Select::from("knowledge_entry")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("acme")),
session: None,
mode: ScopeMode::OwnOrNull,
keys: TableKeys::Uniform,
})
.unwrap();
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT body FROM knowledge_entry WHERE (tenant_id = ?1 OR tenant_id IS NULL) \
AND key_name = ?2 ORDER BY (CASE WHEN tenant_id IS NOT NULL AND tenant_id = ?3 \
THEN ?4 ELSE ?5 END) DESC LIMIT 1"
);
assert_eq!(
params,
vec![
t("acme"),
t("k"),
t("acme"),
SqlValue::Integer(1),
SqlValue::Integer(0)
]
);
}
#[test]
fn is_own_in_select_under_all_uses_the_resolved_own_value() {
let mut q = Select {
columns: vec![item(Expr::IsOwn)],
..Select::from("t")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("acme")),
session: None,
mode: ScopeMode::All,
keys: TableKeys::Uniform,
})
.unwrap();
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT (CASE WHEN tenant_id IS NOT NULL AND tenant_id = ?1 THEN ?2 ELSE ?3 END) FROM t"
);
assert_eq!(
params,
vec![t("acme"), SqlValue::Integer(1), SqlValue::Integer(0)]
);
}
#[test]
fn is_own_without_a_scope_is_rejected() {
let q = Select {
order: vec![OrderBy {
expr: Expr::IsOwn,
dir: Direction::Desc,
}],
..Select::from("t")
};
let err = q.compile(Dialect::Sqlite).unwrap_err();
assert!(
matches!(err, OrmError::BadExpr(m) if m.contains("is_own")),
"expected a fail-closed is_own error, got {err:?}"
);
}
#[test]
fn is_own_in_a_filter_does_not_subtract_the_scope_predicate() {
let mut q = Select {
filter: Some(Predicate::Cmp {
left: Expr::IsOwn,
op: CmpOp::Eq,
right: Expr::Value(SqlValue::Integer(1)),
}),
..Select::from("notes")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("acme")),
session: None,
mode: ScopeMode::OwnOrNull,
keys: TableKeys::Uniform,
})
.unwrap();
let (sql, _) = q.compile(Dialect::Sqlite).unwrap();
assert!(
sql.contains("(tenant_id = ?1 OR tenant_id IS NULL) AND"),
"scope predicate must survive the is_own filter: {sql}"
);
assert!(
sql.contains("CASE WHEN tenant_id IS NOT NULL AND tenant_id = ?2 THEN"),
"is_own lowered to the own-rank CASE: {sql}"
);
}
fn scoped_select(mode: ScopeMode) -> Select {
Select {
filter: Some(cmp("kind", CmpOp::Eq, t("supplier"))),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode,
keys: TableKeys::Uniform,
}),
..Select::from("party")
}
}
#[test]
fn scope_mode_own_or_null_admits_the_shared_baseline() {
let (sql, params) = scoped_select(ScopeMode::OwnOrNull)
.compile(Dialect::Sqlite)
.unwrap();
assert_eq!(
sql,
"SELECT * FROM party WHERE (tenant_id = ?1 OR tenant_id IS NULL) AND kind = ?2"
);
assert_eq!(params, vec![t("ten_1"), t("supplier")]);
}
#[test]
fn scope_mode_null_only_sees_only_the_baseline() {
let (sql, params) = scoped_select(ScopeMode::NullOnly)
.compile(Dialect::Sqlite)
.unwrap();
assert_eq!(
sql,
"SELECT * FROM party WHERE tenant_id IS NULL AND kind = ?1"
);
assert_eq!(params, vec![t("supplier")]);
}
#[test]
fn scope_mode_all_injects_no_tenant_predicate() {
let (sql, params) = scoped_select(ScopeMode::All)
.compile(Dialect::Sqlite)
.unwrap();
assert_eq!(sql, "SELECT * FROM party WHERE kind = ?1");
assert_eq!(params, vec![t("supplier")]);
}
#[test]
fn force_scope_reaches_every_union_branch() {
let branch = Select::from("archived_party");
let mut q = Select {
union: Some(Box::new(Union {
all: false,
query: branch,
})),
..Select::from("party")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
})
.unwrap();
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT * FROM party WHERE tenant_id = ?1 UNION SELECT * FROM archived_party WHERE tenant_id = ?2"
);
assert_eq!(params, vec![t("ten_1"), t("ten_1")]);
}
#[test]
fn scoped_select_scopes_every_joined_table() {
let mut q = Select {
table: "orders".into(),
table_alias: Some("o".into()),
columns: vec![item(Expr::col("v.secret"))],
joins: vec![Join {
kind: JoinKind::Left,
table: "victim".into(),
alias: Some("v".into()),
on: Predicate::Cmp {
left: Expr::col("v.order_id"),
op: CmpOp::Eq,
right: Expr::col("o.id"),
},
}],
..Select::from("orders")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
})
.unwrap();
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT v.secret FROM orders AS o LEFT JOIN victim AS v ON v.order_id = o.id \
WHERE o.tenant_id = ?1 AND v.tenant_id = ?2"
);
assert_eq!(params, vec![t("ten_1"), t("ten_1")]);
}
#[test]
fn scoped_returning_and_distinct_on_subqueries_are_scoped() {
let sub = || Expr::RelatedScalar {
column: "balance".into(),
table: "victim".into(),
filter: Box::new(Predicate::And(Vec::new())),
};
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
};
let mut del = Delete {
table: "orders".into(),
filter: cmp("id", CmpOp::Eq, t("o_1")),
scope: None,
returning: vec![item(sub())],
};
del.force_scope(&scope).unwrap();
let (sql, _) = del.compile(Dialect::Sqlite).unwrap();
assert!(
sql.contains("RETURNING (SELECT balance FROM victim WHERE victim.tenant_id = ?"),
"RETURNING subquery unscoped: {sql}"
);
let mut sel = Select {
columns: vec![item(Expr::col("id"))],
distinct_on: vec![sub()],
..Select::from("orders")
};
sel.force_scope(&scope).unwrap();
let (sql, _) = sel.compile(Dialect::Postgres).unwrap();
assert!(
sql.contains("DISTINCT ON ((SELECT balance FROM victim WHERE victim.tenant_id = ?"),
"DISTINCT ON subquery unscoped: {sql}"
);
}
#[test]
fn scoped_select_scopes_a_subquerys_inner_table() {
let mut q = Select {
columns: vec![item(Expr::RelatedScalar {
column: "balance".into(),
table: "victim".into(),
filter: Box::new(Predicate::And(Vec::new())), })],
..Select::from("orders")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
})
.unwrap();
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT (SELECT balance FROM victim WHERE victim.tenant_id = ?1) \
FROM orders WHERE tenant_id = ?2"
);
assert_eq!(params, vec![t("ten_1"), t("ten_1")]);
}
#[test]
fn insert_select_cannot_forge_the_target_tenant() {
let source = Select {
columns: vec![
item(Expr::val(t("VICTIM"))), item(Expr::col("total")),
],
..Select::from("orders")
};
let mut ins = Insert {
table: "orders".into(),
rows: vec![],
conflict: None,
scope: None,
returning: vec![],
from_select: Some((vec!["TENANT_ID".into(), "total".into()], Box::new(source))),
};
let own = Scope {
column: "tenant_id".into(),
value: Some(t("OWN")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
};
ins.force_scope(Some(&own), Some(&own)).unwrap();
let (sql, params) = ins.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"INSERT INTO orders (total, tenant_id) SELECT total, ?1 FROM orders WHERE tenant_id = ?2"
);
assert_eq!(params, vec![t("OWN"), t("OWN")]);
assert!(
!params.contains(&t("VICTIM")),
"the forged tenant never binds"
);
}
#[test]
fn scoped_update_cannot_reassign_the_tenant() {
let q = Update {
table: "orders".into(),
set: vec![
Assignment {
column: "TENANT_ID".into(),
value: Expr::val(t("VICTIM")),
},
Assignment {
column: "status".into(),
value: Expr::val(t("paid")),
},
],
filter: cmp("id", CmpOp::Eq, t("o_1")),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("OWN")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
returning: vec![],
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"UPDATE orders SET status = ?1 WHERE tenant_id = ?2 AND id = ?3"
);
assert_eq!(params, vec![t("paid"), t("OWN"), t("o_1")]);
assert!(!params.contains(&t("VICTIM")));
}
#[test]
fn scoped_upsert_drops_tenant_reassignment_and_bounds_the_do_update() {
let mut ins = Insert {
table: "orders".into(),
rows: vec![RowValues {
cells: vec![Assignment {
column: "id".into(),
value: Expr::val(t("k")),
}],
}],
conflict: Some(OnConflict {
conflict_columns: vec!["id".into()],
update: vec![
Assignment {
column: "TENANT_ID".into(),
value: Expr::val(t("VICTIM")),
},
Assignment {
column: "total".into(),
value: Expr::val(SqlValue::Integer(999)),
},
],
}),
scope: None,
returning: vec![],
from_select: None,
};
let own = Scope {
column: "tenant_id".into(),
value: Some(t("OWN")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
};
ins.force_scope(Some(&own), Some(&own)).unwrap();
let (sql, params) = ins.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"INSERT INTO orders (id, tenant_id) VALUES (?1, ?2) \
ON CONFLICT (id) DO UPDATE SET total = ?3 WHERE tenant_id = ?4"
);
assert_eq!(
params,
vec![t("k"), t("OWN"), SqlValue::Integer(999), t("OWN")]
);
assert!(!params.contains(&t("VICTIM")));
assert!(matches!(
ins.compile(Dialect::Mysql),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn insert_null_mode_stamps_null_all_mode_stamps_nothing() {
let base = |mode| Insert {
table: "audit_event".into(),
rows: vec![RowValues {
cells: vec![Assignment {
column: "detail".into(),
value: Expr::val(t("x")),
}],
}],
conflict: None,
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode,
keys: TableKeys::Uniform,
}),
returning: vec![],
from_select: None,
};
let (sql, params) = base(ScopeMode::NullOnly).compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"INSERT INTO audit_event (detail, tenant_id) VALUES (?1, ?2)"
);
assert_eq!(params, vec![t("x"), SqlValue::Null]);
let (sql, params) = base(ScopeMode::All).compile(Dialect::Sqlite).unwrap();
assert_eq!(sql, "INSERT INTO audit_event (detail) VALUES (?1)");
assert_eq!(params, vec![t("x")]);
}
#[test]
fn nested_and_or_not_is_parenthesized() {
let q = Select {
filter: Some(all([
Predicate::In {
expr: Expr::col("state"),
values: vec![Expr::val(t("po_linked")), Expr::val(t("awarded"))],
negated: false,
},
any([
cmp("priority", CmpOp::Ge, SqlValue::Integer(3)),
cmp("escalated", CmpOp::Eq, SqlValue::Boolean(true)),
]),
Predicate::Not(Box::new(cmp(
"archived",
CmpOp::Eq,
SqlValue::Boolean(true),
))),
])),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
..Select::from("order_to_network")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT * FROM order_to_network WHERE tenant_id = ?1 AND (state IN (?2, ?3) AND (priority >= ?4 OR escalated = ?5) AND NOT archived = ?6)"
);
assert_eq!(
params,
vec![
t("ten_1"),
t("po_linked"),
t("awarded"),
SqlValue::Integer(3),
SqlValue::Boolean(true),
SqlValue::Boolean(true)
]
);
}
#[test]
fn group_by_having_with_aggregate_and_alias() {
let q = Select {
columns: vec![
item(Expr::col("network_id")),
SelectItem {
expr: Expr::Aggregate(Agg::Sum, Box::new(Expr::col("committed_minor"))),
alias: Some("total".into()),
},
],
group_by: vec![Expr::col("network_id")],
having: Some(Predicate::Cmp {
left: Expr::Aggregate(Agg::Sum, Box::new(Expr::col("committed_minor"))),
op: CmpOp::Gt,
right: Expr::val(SqlValue::Integer(1000)),
}),
order: vec![OrderBy {
expr: Expr::col("total"),
dir: Direction::Desc,
}],
..Select::from("order_to_network")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT network_id, sum(committed_minor) AS total FROM order_to_network GROUP BY network_id HAVING sum(committed_minor) > ?1 ORDER BY total DESC"
);
assert_eq!(params, vec![SqlValue::Integer(1000)]);
}
#[test]
fn join_with_alias_and_column_ref_condition() {
let q = Select {
columns: vec![item(Expr::Aggregate(Agg::Count, Box::new(Expr::Star)))],
joins: vec![Join {
kind: JoinKind::Inner,
table: "element".into(),
alias: Some("e".into()),
on: Predicate::Cmp {
left: Expr::col("order_to_network.element_id"),
op: CmpOp::Eq,
right: Expr::col("e.id"),
},
}],
filter: Some(cmp("order_id", CmpOp::Eq, SqlValue::Integer(7))),
..Select::from("order_to_network")
};
let (sql, _) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT count(*) FROM order_to_network JOIN element AS e ON order_to_network.element_id = e.id WHERE order_id = ?1"
);
}
#[test]
fn between_like_insensitive_and_notin() {
let q = Select {
filter: Some(all([
Predicate::Between {
expr: Expr::col("amount"),
low: Expr::val(SqlValue::Integer(10)),
high: Expr::val(SqlValue::Integer(20)),
negated: false,
},
Predicate::Like {
expr: Expr::col("name"),
pattern: "ac%".into(),
insensitive: true,
negated: false,
},
Predicate::In {
expr: Expr::col("state"),
values: vec![Expr::val(t("void"))],
negated: true,
},
])),
..Select::from("invoice")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT * FROM invoice WHERE amount BETWEEN ?1 AND ?2 AND lower(name) LIKE lower(?3) AND state NOT IN (?4)"
);
assert_eq!(
params,
vec![
SqlValue::Integer(10),
SqlValue::Integer(20),
t("ac%"),
t("void")
]
);
}
#[test]
fn arithmetic_and_functions_in_select_and_set() {
let q = Select {
columns: vec![
SelectItem {
expr: Expr::Func(Func::Lower, vec![Expr::col("email")]),
alias: Some("email_lc".into()),
},
item(Expr::Binary(
BinOp::Mul,
Box::new(Expr::col("qty")),
Box::new(Expr::val(SqlValue::Integer(2))),
)),
item(Expr::Func(
Func::Coalesce,
vec![Expr::col("nickname"), Expr::val(t("n/a"))],
)),
],
..Select::from("account")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT lower(email) AS email_lc, (qty * ?1), coalesce(nickname, ?2) FROM account"
);
assert_eq!(params, vec![SqlValue::Integer(2), t("n/a")]);
}
#[test]
fn empty_in_and_not_in_are_identities() {
let matches_none = Select {
filter: Some(Predicate::In {
expr: Expr::col("x"),
values: vec![],
negated: false,
}),
..Select::from("t")
};
assert_eq!(
matches_none.compile(Dialect::Sqlite).unwrap().0,
"SELECT * FROM t WHERE 1 = 0"
);
let matches_all = Select {
filter: Some(Predicate::In {
expr: Expr::col("x"),
values: vec![],
negated: true,
}),
..Select::from("t")
};
assert_eq!(
matches_all.compile(Dialect::Sqlite).unwrap().0,
"SELECT * FROM t WHERE 1 = 1"
);
}
#[test]
fn insert_with_scope_and_returning() {
let q = Insert {
table: "work_area".into(),
rows: vec![RowValues {
cells: vec![
Assignment {
column: "id".into(),
value: Expr::val(t("wa_1")),
},
Assignment {
column: "project_id".into(),
value: Expr::val(t("prj_1")),
},
],
}],
conflict: None,
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
returning: vec![item(Expr::col("id"))],
from_select: None,
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"INSERT INTO work_area (id, project_id, tenant_id) VALUES (?1, ?2, ?3) RETURNING id"
);
assert_eq!(params, vec![t("wa_1"), t("prj_1"), t("ten_1")]);
}
#[test]
fn upsert_do_update_and_do_nothing() {
let base = |update: Vec<Assignment>| Insert {
table: "country_pack".into(),
rows: vec![RowValues {
cells: vec![
Assignment {
column: "country".into(),
value: Expr::val(t("US")),
},
Assignment {
column: "currency".into(),
value: Expr::val(t("USD")),
},
],
}],
conflict: Some(OnConflict {
conflict_columns: vec!["tenant_id".into(), "country".into()],
update,
}),
scope: None,
returning: vec![],
from_select: None,
};
let (sql_do, _) = base(vec![Assignment {
column: "currency".into(),
value: Expr::val(t("USD")),
}])
.compile(Dialect::Sqlite)
.unwrap();
assert_eq!(
sql_do,
"INSERT INTO country_pack (country, currency) VALUES (?1, ?2) ON CONFLICT (tenant_id, country) DO UPDATE SET currency = ?3"
);
let (sql_nothing, _) = base(vec![]).compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql_nothing,
"INSERT INTO country_pack (country, currency) VALUES (?1, ?2) ON CONFLICT (tenant_id, country) DO NOTHING"
);
}
#[test]
fn update_binds_set_before_where_and_supports_expr_set() {
let q = Update {
table: "counter".into(),
set: vec![Assignment {
column: "hits".into(),
value: Expr::Binary(
BinOp::Add,
Box::new(Expr::col("hits")),
Box::new(Expr::val(SqlValue::Integer(1))),
),
}],
filter: cmp("id", CmpOp::Eq, t("c_1")),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
returning: vec![],
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"UPDATE counter SET hits = (hits + ?1) WHERE tenant_id = ?2 AND id = ?3"
);
assert_eq!(params, vec![SqlValue::Integer(1), t("ten_1"), t("c_1")]);
}
#[test]
fn identifier_injection_is_rejected() {
let q = Select {
columns: vec![item(Expr::col("id; DROP TABLE users"))],
..Select::from("t")
};
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::InvalidIdentifier(_))
));
}
#[test]
fn qualified_identifier_allowed() {
let q = Select {
columns: vec![item(Expr::col("t.id"))],
..Select::from("t")
};
assert_eq!(q.compile(Dialect::Sqlite).unwrap().0, "SELECT t.id FROM t");
}
#[test]
fn function_arity_is_checked() {
let q = Select {
columns: vec![item(Expr::Func(Func::Lower, vec![]))],
..Select::from("t")
};
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn update_with_empty_all_filter_is_refused() {
let q = Update {
table: "t".into(),
set: vec![Assignment {
column: "x".into(),
value: Expr::val(SqlValue::Integer(1)),
}],
filter: Predicate::And(vec![]),
scope: None,
returning: vec![],
};
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::Empty(_))
));
}
#[test]
fn empty_filter_with_scope_is_allowed() {
let q = Update {
table: "t".into(),
set: vec![Assignment {
column: "x".into(),
value: Expr::val(SqlValue::Integer(1)),
}],
filter: Predicate::And(vec![]),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
returning: vec![],
};
assert_eq!(
q.compile(Dialect::Sqlite).unwrap().0,
"UPDATE t SET x = ?1 WHERE tenant_id = ?2"
);
}
#[test]
fn delete_by_predicate_compiles() {
let q = Delete {
table: "payment".into(),
filter: cmp("id", CmpOp::Eq, t("pay_1")),
scope: None,
returning: vec![],
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(sql, "DELETE FROM payment WHERE id = ?1");
assert_eq!(params, vec![t("pay_1")]);
}
#[test]
fn delete_returning_renders() {
let q = Delete {
table: "pending_signup".into(),
filter: cmp("slug", CmpOp::Eq, t("acme")),
scope: None,
returning: vec![item(Expr::col("name")), item(Expr::col("password_hash"))],
};
assert_eq!(
q.compile(Dialect::Postgres).unwrap().0,
"DELETE FROM pending_signup WHERE slug = ?1 RETURNING name, password_hash"
);
}
#[test]
fn delete_with_empty_filter_is_refused() {
let q = Delete {
table: "t".into(),
filter: Predicate::And(vec![]),
scope: None,
returning: vec![],
};
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::Empty(_))
));
}
#[test]
fn delete_empty_filter_with_scope_is_allowed() {
let q = Delete {
table: "t".into(),
filter: Predicate::And(vec![]),
scope: Some(Scope {
column: "tenant_id".into(),
value: Some(t("ten_1")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::Uniform,
}),
returning: vec![],
};
assert_eq!(
q.compile(Dialect::Sqlite).unwrap().0,
"DELETE FROM t WHERE tenant_id = ?1"
);
}
#[test]
fn delete_rejects_identifier_injection_in_table() {
let q = Delete {
table: "t; DROP TABLE users".into(),
filter: cmp("id", CmpOp::Eq, t("x")),
scope: None,
returning: vec![],
};
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::InvalidIdentifier(_))
));
}
#[test]
fn case_expression_renders_with_bound_params() {
let q = Select {
columns: vec![item(Expr::Case {
branches: vec![(
cmp("state", CmpOp::Eq, t("open")),
Expr::val(SqlValue::Integer(1)),
)],
otherwise: Some(Box::new(Expr::val(SqlValue::Integer(0)))),
})],
..Select::from("t")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT (CASE WHEN state = ?1 THEN ?2 ELSE ?3 END) FROM t"
);
assert_eq!(
params,
vec![t("open"), SqlValue::Integer(1), SqlValue::Integer(0)]
);
}
#[test]
fn distinct_on_renders_on_postgres_and_fails_closed_elsewhere() {
let q = Select {
distinct_on: vec![Expr::col("key")],
columns: vec![item(Expr::col("key")), item(Expr::col("val"))],
..Select::from("consent_state")
};
assert_eq!(
q.compile(Dialect::Postgres).unwrap().0,
"SELECT DISTINCT ON (key) key, val FROM consent_state"
);
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn empty_case_is_rejected() {
let q = Select {
columns: vec![item(Expr::Case {
branches: vec![],
otherwise: None,
})],
..Select::from("t")
};
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn json_extract_dyn_binds_the_key() {
let q = Select {
columns: vec![item(Expr::JsonExtractDyn(
Box::new(Expr::col("labels")),
Box::new(Expr::val(t("en"))),
))],
..Select::from("vocabulary_term")
};
for d in [Dialect::Postgres, Dialect::Sqlite] {
assert_eq!(
q.compile(d).unwrap().0,
"SELECT (labels ->> ?1) FROM vocabulary_term"
);
}
assert!(matches!(
q.compile(Dialect::Mysql),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn json_concat_merge_is_postgres_only() {
let q = Update {
table: "request".into(),
set: vec![Assignment {
column: "brief_state".into(),
value: Expr::JsonConcat(
Box::new(Expr::col("brief_state")),
Box::new(Expr::val(SqlValue::Json("{\"a\":1}".into()))),
),
}],
filter: cmp("id", CmpOp::Eq, t("req_1")),
scope: None,
returning: vec![],
};
assert_eq!(
q.compile(Dialect::Postgres).unwrap().0,
"UPDATE request SET brief_state = (brief_state || ?1) WHERE id = ?2"
);
assert!(matches!(
q.compile(Dialect::Sqlite),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn union_renders_both_bodies_with_shared_params() {
let q = Select {
columns: vec![item(Expr::col("slug"))],
filter: Some(cmp("slug", CmpOp::Eq, t("acme"))),
union: Some(Box::new(Union {
all: false,
query: Select {
columns: vec![item(Expr::col("slug"))],
filter: Some(cmp("slug", CmpOp::Eq, t("acme"))),
..Select::from("reserved_slug")
},
})),
..Select::from("pending_signup")
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT slug FROM pending_signup WHERE slug = ?1 \
UNION SELECT slug FROM reserved_slug WHERE slug = ?2"
);
assert_eq!(params, vec![t("acme"), t("acme")]);
}
#[test]
fn insert_from_select_shares_params_and_carries_no_auto_scope() {
let q = Insert {
table: "portfolio_ref".into(),
rows: vec![],
conflict: None,
scope: None,
returning: vec![],
from_select: Some((
vec!["a".into(), "b".into()],
Box::new(Select {
columns: vec![item(Expr::col("x")), item(Expr::col("y"))],
filter: Some(cmp("id", CmpOp::Eq, t("pi_1"))),
..Select::from("portfolio_item")
}),
)),
};
let (sql, params) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"INSERT INTO portfolio_ref (a, b) SELECT x, y FROM portfolio_item WHERE id = ?1"
);
assert_eq!(params, vec![t("pi_1")]);
}
#[test]
fn related_scalar_and_in_subquery_render() {
let q = Select {
columns: vec![item(Expr::col("id"))],
filter: Some(Predicate::Cmp {
left: Expr::col("id"),
op: CmpOp::Eq,
right: Expr::RelatedScalar {
column: "head_version".into(),
table: "pack".into(),
filter: Box::new(cmp("id", CmpOp::Eq, t("pk_1"))),
},
}),
..Select::from("pack_version")
};
assert_eq!(
q.compile(Dialect::Sqlite).unwrap().0,
"SELECT id FROM pack_version WHERE id = (SELECT head_version FROM pack WHERE id = ?1)"
);
let q2 = Select {
columns: vec![item(Expr::col("x"))],
filter: Some(Predicate::InSubquery {
expr: Expr::col("doc_id"),
column: "id".into(),
table: "document".into(),
filter: Box::new(cmp("tenant_id", CmpOp::Eq, t("ten_1"))),
negated: false,
}),
..Select::from("access")
};
assert_eq!(
q2.compile(Dialect::Sqlite).unwrap().0,
"SELECT x FROM access WHERE doc_id IN (SELECT id FROM document WHERE tenant_id = ?1)"
);
}
#[test]
fn now_renders_without_parens() {
let q = Select {
columns: vec![item(Expr::Func(Func::Now, vec![]))],
..Select::from("t")
};
assert_eq!(
q.compile(Dialect::Sqlite).unwrap().0,
"SELECT current_timestamp FROM t"
);
}
fn json_query() -> Select {
Select {
columns: vec![item(Expr::JsonExtract(
Box::new(Expr::col("metadata")),
vec!["status".into()],
))],
filter: Some(Predicate::Cmp {
left: Expr::JsonExtract(
Box::new(Expr::col("metadata")),
vec!["a".into(), "b".into()],
),
op: CmpOp::Eq,
right: Expr::val(t("x")),
}),
..Select::from("doc")
}
}
#[test]
fn json_extract_sqlite_and_mysql_bind_the_path() {
for d in [Dialect::Sqlite, Dialect::Mysql] {
let (sql, params) = json_query().compile(d).unwrap();
assert_eq!(
sql,
"SELECT json_extract(metadata, ?1) FROM doc WHERE json_extract(metadata, ?2) = ?3"
);
assert_eq!(params, vec![t("$.status"), t("$.a.b"), t("x")]);
}
}
#[test]
fn json_extract_postgres_inlines_the_validated_path() {
let (sql, params) = json_query().compile(Dialect::Postgres).unwrap();
assert_eq!(
sql,
"SELECT (metadata) #>> '{status}' FROM doc WHERE (metadata) #>> '{a,b}' = ?1"
);
assert_eq!(params, vec![t("x")]);
}
#[test]
fn json_extract_key_injection_is_rejected() {
let q = Select {
columns: vec![item(Expr::JsonExtract(
Box::new(Expr::col("m")),
vec!["a'); DROP TABLE t--".into()],
))],
..Select::from("doc")
};
assert!(matches!(
q.compile(Dialect::Postgres),
Err(OrmError::InvalidIdentifier(_))
));
}
fn knn_query() -> Select {
Select {
columns: vec![item(Expr::col("id"))],
order: vec![OrderBy {
expr: Expr::Distance {
left: Box::new(Expr::col("embedding")),
right: Box::new(Expr::VectorLiteral("[0.1, 0.2, 0.3]".into())),
metric: Metric::Cosine,
},
dir: Direction::Asc,
}],
limit: Some(5),
..Select::from("doc")
}
}
#[test]
fn distance_orders_by_cosine_nearest_neighbour_on_postgres() {
let (sql, params) = knn_query().compile(Dialect::Postgres).unwrap();
assert_eq!(
sql,
"SELECT id FROM doc ORDER BY (embedding <=> ?1::vector) ASC LIMIT 5"
);
assert_eq!(params, vec![t("[0.1,0.2,0.3]")]);
}
#[test]
fn distance_l2_in_select_list_on_postgres() {
let q = Select {
columns: vec![
item(Expr::col("id")),
SelectItem {
expr: Expr::Distance {
left: Box::new(Expr::col("embedding")),
right: Box::new(Expr::VectorLiteral("[-1, 2e0, 3.5]".into())),
metric: Metric::L2,
},
alias: Some("dist".into()),
},
],
..Select::from("doc")
};
let (sql, params) = q.compile(Dialect::Postgres).unwrap();
assert_eq!(
sql,
"SELECT id, (embedding <-> ?1::vector) AS dist FROM doc"
);
assert_eq!(params, vec![t("[-1,2e0,3.5]")]);
}
#[test]
fn distance_fails_closed_off_postgres() {
for d in [Dialect::Sqlite, Dialect::Mysql] {
assert!(
matches!(knn_query().compile(d), Err(OrmError::BadExpr(_))),
"vector distance must be rejected on {d:?}"
);
}
}
#[test]
fn vector_literal_fails_closed_off_postgres() {
for d in [Dialect::Sqlite, Dialect::Mysql] {
let q = Select {
columns: vec![item(Expr::VectorLiteral("[1, 2]".into()))],
..Select::from("doc")
};
assert!(
matches!(q.compile(d), Err(OrmError::BadExpr(_))),
"vector literal must be rejected on {d:?}"
);
}
}
#[test]
fn malformed_vector_literal_is_rejected() {
for bad in [
"1,2",
"[a, b]",
"[]",
"[1, 2",
"[1,,2]",
"[Infinity]",
"[1, NaN]",
] {
let q = Select {
columns: vec![item(Expr::VectorLiteral(bad.to_string()))],
..Select::from("doc")
};
assert!(
matches!(q.compile(Dialect::Postgres), Err(OrmError::BadExpr(_))),
"expected {bad:?} to be rejected"
);
}
}
fn related(agg: Agg, arg: RelArg, table: &str, filter: Predicate) -> Expr {
Expr::RelatedAggregate {
agg,
arg,
table: table.into(),
filter: Box::new(filter),
}
}
fn correlate(fk: &str, pk: &str) -> Predicate {
Predicate::Cmp {
left: Expr::col(fk),
op: CmpOp::Eq,
right: Expr::col(pk),
}
}
#[test]
fn related_aggregate_single_correlated_count() {
let q = Select {
columns: vec![
item(Expr::col("id")),
SelectItem {
expr: related(
Agg::Count,
RelArg::Star,
"element",
correlate("element.order_id", "work_order.id"),
),
alias: Some("element_count".into()),
},
],
..Select::from("work_order")
};
let (sql, params) = q.compile(Dialect::Postgres).unwrap();
assert_eq!(
sql,
"SELECT id, (SELECT count(*) FROM element WHERE element.order_id = work_order.id) AS element_count FROM work_order"
);
assert!(params.is_empty());
}
#[test]
fn related_aggregate_two_counts_bind_distinct_params_and_dont_fan_out() {
let with_status = |child: &str, fk: &str, status: &str| {
related(
Agg::Count,
RelArg::Star,
child,
Predicate::And(vec![
correlate(fk, "party.id"),
Predicate::Cmp {
left: Expr::col("status"),
op: CmpOp::Eq,
right: Expr::val(t(status)),
},
]),
)
};
let q = Select {
columns: vec![
item(with_status("party_role", "party_role.party_id", "active")),
item(with_status(
"party_qualification",
"party_qualification.party_id",
"valid",
)),
],
..Select::from("party")
};
let (sql, params) = q.compile(Dialect::Postgres).unwrap();
assert_eq!(
sql,
"SELECT \
(SELECT count(*) FROM party_role WHERE party_role.party_id = party.id AND status = ?1), \
(SELECT count(*) FROM party_qualification WHERE party_qualification.party_id = party.id AND status = ?2) \
FROM party"
);
assert_eq!(params, vec![t("active"), t("valid")]);
}
#[test]
fn related_aggregate_with_temporal_or_filter() {
let q = Select {
columns: vec![SelectItem {
expr: related(
Agg::Count,
RelArg::Star,
"party_qualification",
Predicate::And(vec![
correlate("party_qualification.party_id", "party.id"),
Predicate::Or(vec![
Predicate::Null {
expr: Expr::col("valid_to"),
negated: false,
},
Predicate::Cmp {
left: Expr::col("valid_to"),
op: CmpOp::Gt,
right: Expr::val(t("2026-01-01")),
},
]),
]),
),
alias: Some("active_quals".into()),
}],
..Select::from("party")
};
let (sql, params) = q.compile(Dialect::Postgres).unwrap();
assert_eq!(
sql,
"SELECT (SELECT count(*) FROM party_qualification WHERE party_qualification.party_id = party.id AND (valid_to IS NULL OR valid_to > ?1)) AS active_quals FROM party"
);
assert_eq!(params, vec![t("2026-01-01")]);
}
#[test]
fn related_aggregate_max_over_a_column_is_portable() {
let q = Select {
columns: vec![SelectItem {
expr: related(
Agg::Max,
RelArg::Column("total_minor".into()),
"line_item",
correlate("line_item.order_id", "order_summary.id"),
),
alias: Some("max_total".into()),
}],
..Select::from("order_summary")
};
let (sql, _) = q.compile(Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"SELECT (SELECT max(total_minor) FROM line_item WHERE line_item.order_id = order_summary.id) AS max_total FROM order_summary"
);
}
#[test]
fn related_aggregate_star_is_count_only() {
let q = Select {
columns: vec![item(related(
Agg::Sum,
RelArg::Star,
"t",
correlate("t.fk", "p.id"),
))],
..Select::from("p")
};
assert!(matches!(
q.compile(Dialect::Postgres),
Err(OrmError::BadExpr(_))
));
}
#[test]
fn related_aggregate_table_injection_is_rejected() {
let q = Select {
columns: vec![item(related(
Agg::Count,
RelArg::Star,
"element; DROP TABLE users",
correlate("element.order_id", "p.id"),
))],
..Select::from("p")
};
assert!(matches!(
q.compile(Dialect::Postgres),
Err(OrmError::InvalidIdentifier(_))
));
}
#[test]
fn target_read_conjoins_the_public_subset_and_composes_across_joins() {
use std::collections::BTreeMap;
let keys = BTreeMap::from([
(
"products".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
),
(
"reviews".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
),
]);
let public = BTreeMap::from([
(
"products".to_string(),
vec![
PublicTermSql::Cmp {
column: "published".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
},
PublicTermSql::Null {
column: "deleted_at".into(),
negated: false,
},
],
),
(
"reviews".to_string(),
vec![PublicTermSql::Cmp {
column: "visible".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
}],
),
]);
let mut q = Select {
table_alias: Some("p".into()),
joins: vec![Join {
kind: JoinKind::Left,
table: "reviews".into(),
alias: Some("r".into()),
on: Predicate::Cmp {
left: Expr::col("p.id"),
op: CmpOp::Eq,
right: Expr::col("r.product_id"),
},
}],
..Select::from("products")
};
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("B")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTableTarget {
keys: keys.clone(),
public: public.clone(),
write: std::collections::BTreeSet::new(),
require_public: true,
},
})
.unwrap();
let (sql, _params) = q.compile(Dialect::Sqlite).unwrap();
assert!(sql.contains("p.tenant_id = ?"), "base tenant scope: {sql}");
assert!(sql.contains("p.published = ?"), "base public term: {sql}");
assert!(
sql.contains("p.deleted_at IS NULL"),
"base public null term: {sql}"
);
assert!(
sql.contains("r.tenant_id = ?"),
"joined tenant scope: {sql}"
);
assert!(sql.contains("r.visible = ?"), "joined public term: {sql}");
}
#[test]
fn target_read_of_a_table_with_no_public_subset_is_refused() {
use std::collections::BTreeMap;
let keys = BTreeMap::from([(
"secret_table".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
)]);
let deny = Scope {
column: "tenant_id".into(),
value: Some(t("B")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTableTarget {
keys,
public: BTreeMap::new(), write: std::collections::BTreeSet::new(),
require_public: true,
},
};
let mut q = Select::from("secret_table");
let err = q
.force_scope(&deny)
.err()
.or_else(|| q.compile(Dialect::Sqlite).err());
assert!(
matches!(&err, Some(OrmError::PublicSubsetUndeclared(t)) if t == "secret_table"),
"expected PublicSubsetUndeclared, got {err:?}"
);
}
#[test]
fn own_read_is_unaffected_by_the_public_injection() {
use std::collections::BTreeMap;
let mut q = Select::from("products");
q.force_scope(&Scope {
column: "tenant_id".into(),
value: Some(t("A")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([(
"products".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
)])),
})
.unwrap();
let (sql, _p) = q.compile(Dialect::Sqlite).unwrap();
assert!(sql.contains("tenant_id = ?"));
assert!(
!sql.contains("published") && !sql.contains("IS NULL"),
"own read must carry no public confinement: {sql}"
);
}
fn target_write_scope(write: &[&str]) -> Scope {
use std::collections::{BTreeMap, BTreeSet};
Scope {
column: "tenant_id".into(),
value: Some(t("tenant_B")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTableTarget {
keys: BTreeMap::from([(
"products".to_string(),
ResolvedScope::Column("tenant_id".to_string()),
)]),
public: BTreeMap::from([(
"products".to_string(),
vec![
PublicTermSql::Cmp {
column: "published".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
},
PublicTermSql::Null {
column: "deleted_at".into(),
negated: false,
},
],
)]),
write: write
.iter()
.map(ToString::to_string)
.collect::<BTreeSet<_>>(),
require_public: true,
},
}
}
fn target_insert(cells: Vec<Assignment>) -> Insert {
Insert {
table: "products".into(),
rows: vec![RowValues { cells }],
conflict: None,
scope: None,
returning: vec![],
from_select: None,
}
}
#[test]
fn target_insert_forces_tenant_and_public_and_accepts_only_allowlisted_columns() {
let scope = target_write_scope(&["title"]);
let mut ins = target_insert(vec![Assignment {
column: "title".into(),
value: Expr::val(t("Hello")),
}]);
ins.force_scope(Some(&scope), Some(&scope)).unwrap();
let (sql, params) = ins.compile(Dialect::Sqlite).unwrap();
assert!(sql.contains("tenant_id"), "{sql}");
assert!(sql.contains("published"), "{sql}");
assert!(sql.contains("deleted_at"), "{sql}");
assert!(
params.contains(&t("tenant_B")),
"tenant forced to B: {params:?}"
);
assert!(
params.contains(&SqlValue::Boolean(true)),
"published forced true: {params:?}"
);
assert!(
params.contains(&SqlValue::Null),
"deleted_at forced NULL: {params:?}"
);
assert!(params.contains(&t("Hello")), "guest title kept: {params:?}");
}
#[test]
fn target_insert_refuses_a_non_allowlisted_column() {
let scope = target_write_scope(&["title"]);
let mut ins = target_insert(vec![
Assignment {
column: "title".into(),
value: Expr::val(t("x")),
},
Assignment {
column: "price".into(),
value: Expr::val(SqlValue::Integer(9)),
},
]);
let err = ins.force_scope(Some(&scope), Some(&scope)).unwrap_err();
assert!(
matches!(err, OrmError::TargetWriteColumnDenied(ref c) if c == "price"),
"{err:?}"
);
}
#[test]
fn target_insert_refuses_setting_the_tenant_or_visibility_column() {
for bad in ["tenant_id", "published", "deleted_at"] {
let scope = target_write_scope(&["title", bad]); let mut ins = target_insert(vec![Assignment {
column: bad.into(),
value: Expr::val(t("x")),
}]);
let err = ins.force_scope(Some(&scope), Some(&scope)).unwrap_err();
assert!(
matches!(err, OrmError::TargetWriteColumnDenied(ref c) if c == bad),
"{bad}: {err:?}"
);
}
}
#[test]
fn target_insert_with_no_write_grant_is_refused() {
let scope = target_write_scope(&[]); let mut ins = target_insert(vec![Assignment {
column: "title".into(),
value: Expr::val(t("x")),
}]);
let err = ins.force_scope(Some(&scope), Some(&scope)).unwrap_err();
assert!(
matches!(err, OrmError::TargetWriteNotGranted(ref t) if t == "products"),
"{err:?}"
);
}
#[test]
fn target_insert_select_and_upsert_are_refused() {
let scope = target_write_scope(&["title"]);
let mut ins = target_insert(vec![Assignment {
column: "title".into(),
value: Expr::val(t("x")),
}]);
ins.from_select = Some((vec!["title".into()], Box::new(Select::from("products"))));
assert!(matches!(
ins.force_scope(Some(&scope), Some(&scope)).unwrap_err(),
OrmError::TargetWriteUnsupported("INSERT … SELECT")
));
let mut ins2 = target_insert(vec![Assignment {
column: "title".into(),
value: Expr::val(t("x")),
}]);
ins2.conflict = Some(OnConflict {
conflict_columns: vec![],
update: vec![],
});
assert!(matches!(
ins2.force_scope(Some(&scope), Some(&scope)).unwrap_err(),
OrmError::TargetWriteUnsupported("ON CONFLICT upsert")
));
}
#[test]
fn target_update_confines_to_the_public_subset_and_enforces_the_allowlist() {
let scope = target_write_scope(&["title"]);
let mut upd = Update {
table: "products".into(),
set: vec![Assignment {
column: "title".into(),
value: Expr::val(t("new")),
}],
filter: cmp("id", CmpOp::Eq, t("p1")),
scope: None,
returning: vec![],
};
upd.force_scope(&scope).unwrap();
let (sql, params) = upd.compile(Dialect::Sqlite).unwrap();
assert!(sql.contains("tenant_id = ?"), "tenant confinement: {sql}");
assert!(sql.contains("published = ?"), "public confinement: {sql}");
assert!(
sql.contains("deleted_at IS NULL"),
"public null confinement: {sql}"
);
assert!(sql.contains("SET title = ?"), "{sql}");
assert!(params.contains(&t("tenant_B")), "{params:?}");
}
#[test]
fn target_update_capability_no_subset_confines_tenant_only_not_refused() {
use std::collections::{BTreeMap, BTreeSet};
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("tenant_B")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTableTarget {
keys: BTreeMap::from([(
"invoices".to_string(),
ResolvedScope::Column("tenant_id".into()),
)]),
public: BTreeMap::new(), write: BTreeSet::from(["amount".to_string()]),
require_public: false, },
};
let mut upd = Update {
table: "invoices".into(),
set: vec![Assignment {
column: "amount".into(),
value: Expr::val(SqlValue::Integer(5)),
}],
filter: cmp("id", CmpOp::Eq, t("inv1")),
scope: None,
returning: vec![],
};
upd.force_scope(&scope).unwrap();
let (sql, _params) = upd.compile(Dialect::Sqlite).unwrap();
assert!(
sql.contains("tenant_id = ?"),
"tenant=B confinement present: {sql}"
);
assert!(!sql.contains("published"), "no visibility conjunct: {sql}");
assert!(sql.contains("SET amount = ?"), "{sql}");
let mut bad = Update {
table: "invoices".into(),
set: vec![Assignment {
column: "tenant_id".into(),
value: Expr::val(t("evil")),
}],
filter: cmp("id", CmpOp::Eq, t("inv1")),
scope: None,
returning: vec![],
};
assert!(
bad.force_scope(&scope).is_err(),
"tenant column still un-settable"
);
}
#[test]
fn target_update_refuses_a_non_allowlisted_or_visibility_set() {
for bad in ["price", "published", "tenant_id"] {
let scope = target_write_scope(&["title"]);
let mut upd = Update {
table: "products".into(),
set: vec![Assignment {
column: bad.into(),
value: Expr::val(t("x")),
}],
filter: cmp("id", CmpOp::Eq, t("p1")),
scope: None,
returning: vec![],
};
assert!(
matches!(upd.force_scope(&scope).unwrap_err(), OrmError::TargetWriteColumnDenied(ref c) if c == bad),
"{bad}"
);
}
}
#[test]
fn target_write_refuses_a_qualified_column() {
let scope = target_write_scope(&["title"]);
let mut ins = target_insert(vec![Assignment {
column: "published.x".into(),
value: Expr::val(t("x")),
}]);
assert!(matches!(
ins.force_scope(Some(&scope), Some(&scope)).unwrap_err(),
OrmError::TargetWriteColumnDenied(ref c) if c == "published.x"
));
let mut upd = Update {
table: "products".into(),
set: vec![Assignment {
column: "title.y".into(),
value: Expr::val(t("x")),
}],
filter: cmp("id", CmpOp::Eq, t("p1")),
scope: None,
returning: vec![],
};
assert!(matches!(
upd.force_scope(&scope).unwrap_err(),
OrmError::TargetWriteColumnDenied(ref c) if c == "title.y"
));
}
#[test]
fn target_delete_is_always_refused() {
let scope = target_write_scope(&["title"]);
let mut del = Delete {
table: "products".into(),
filter: cmp("id", CmpOp::Eq, t("p1")),
scope: None,
returning: vec![],
};
assert!(matches!(
del.force_scope(&scope).unwrap_err(),
OrmError::TargetDeleteRefused(t) if t == "products"
));
}
#[test]
fn target_promote_is_refused() {
let scope = target_write_scope(&["title"]);
assert!(matches!(
compile_promote(&scope, "products", Dialect::Sqlite).unwrap_err(),
OrmError::TargetWriteUnsupported("promote")
));
}
fn attach_spec() -> AttachReference {
AttachReference {
child: "favorites".into(),
parent: "products".into(),
ref_column: "id".into(),
ref_value: t("prod_1"),
set: vec![Assignment {
column: "note".into(),
value: Expr::val(t("nice")),
}],
}
}
#[test]
fn attach_reference_own_derives_tenant_from_the_scoped_parent() {
use std::collections::BTreeMap;
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("A")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([
(
"favorites".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
(
"products".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
])),
};
let (sql, params) =
compile_attach_reference(&scope, &attach_spec(), Dialect::Sqlite).unwrap();
assert_eq!(
sql,
"INSERT INTO favorites (note, tenant_id) SELECT ?1, tenant_id FROM products \
WHERE tenant_id = ?2 AND id = ?3"
);
assert_eq!(params, vec![t("nice"), t("A"), t("prod_1")]);
}
#[test]
fn attach_reference_target_confines_parent_to_b_public_and_forces_child_public() {
use std::collections::{BTreeMap, BTreeSet};
let public_terms = vec![PublicTermSql::Cmp {
column: "visible".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
}];
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("tenant_B")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTableTarget {
keys: BTreeMap::from([
(
"favorites".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
(
"products".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
]),
public: BTreeMap::from([
("favorites".to_string(), public_terms.clone()),
(
"products".to_string(),
vec![PublicTermSql::Cmp {
column: "published".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
}],
),
]),
write: BTreeSet::from(["note".to_string()]),
require_public: true,
},
};
let (sql, params) =
compile_attach_reference(&scope, &attach_spec(), Dialect::Sqlite).unwrap();
assert!(
sql.contains("INSERT INTO favorites (note, tenant_id, visible)"),
"{sql}"
);
assert!(
sql.contains("SELECT ?1, tenant_id, ?2 FROM products"),
"{sql}"
);
assert!(
sql.contains("published = ?") && sql.contains("tenant_id = ?"),
"parent confined: {sql}"
);
assert!(sql.contains("AND id = ?"), "ref selector present: {sql}");
assert!(
params.contains(&t("tenant_B")),
"parent confined to B: {params:?}"
);
assert!(
params.contains(&SqlValue::Boolean(true)),
"child visible forced + parent published: {params:?}"
);
}
#[test]
fn attach_reference_target_refuses_a_non_allowlisted_or_visibility_set() {
use std::collections::{BTreeMap, BTreeSet};
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("tenant_B")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTableTarget {
keys: BTreeMap::from([
(
"favorites".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
(
"products".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
]),
public: BTreeMap::from([
(
"favorites".to_string(),
vec![PublicTermSql::Cmp {
column: "visible".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
}],
),
(
"products".to_string(),
vec![PublicTermSql::Cmp {
column: "published".into(),
op: CmpOp::Eq,
value: SqlValue::Boolean(true),
}],
),
]),
write: BTreeSet::from(["note".to_string()]),
require_public: true,
},
};
for bad in ["price", "visible", "tenant_id"] {
let spec = AttachReference {
set: vec![Assignment {
column: bad.into(),
value: Expr::val(t("x")),
}],
..attach_spec()
};
assert!(
matches!(
compile_attach_reference(&scope, &spec, Dialect::Sqlite).unwrap_err(),
OrmError::TargetWriteColumnDenied(ref c) if c == bad
),
"{bad}"
);
}
}
#[test]
fn attach_reference_own_refuses_the_guest_naming_the_tenant_column() {
use std::collections::BTreeMap;
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("A")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([
(
"favorites".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
(
"products".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
])),
};
let spec = AttachReference {
set: vec![Assignment {
column: "tenant_id".into(),
value: Expr::val(t("VICTIM")),
}],
..attach_spec()
};
assert!(matches!(
compile_attach_reference(&scope, &spec, Dialect::Sqlite).unwrap_err(),
OrmError::TargetWriteColumnDenied(ref c) if c == "tenant_id"
));
}
#[test]
fn attach_reference_refuses_a_non_column_table() {
use std::collections::BTreeMap;
let scope = Scope {
column: "tenant_id".into(),
value: Some(t("A")),
session: None,
mode: ScopeMode::Own,
keys: TableKeys::PerTable(BTreeMap::from([
(
"favorites".to_string(),
ResolvedScope::Column("tenant_id".into()),
),
("countries".to_string(), ResolvedScope::Unscoped),
])),
};
let spec = AttachReference {
parent: "countries".into(),
..attach_spec()
};
assert!(compile_attach_reference(&scope, &spec, Dialect::Sqlite).is_err());
}
}