use toasty_core::{
schema::app::{self, BelongsTo, FieldId, FieldTy, ModelId},
stmt::{self, Expr, ExprContext, IntoExprTarget, Visit, VisitMut},
};
use super::relation_expr;
pub(super) struct LiftInSubquery<'a> {
cx: ExprContext<'a>,
exclude_nulls: bool,
}
impl<'a> LiftInSubquery<'a> {
pub(super) fn new(cx: ExprContext<'a>, exclude_nulls: bool) -> Self {
Self { cx, exclude_nulls }
}
pub(super) fn rewrite(&mut self, stmt: &mut stmt::Statement) {
self.visit_mut(stmt);
}
fn scope<'scope>(&'scope self, target: impl IntoExprTarget<'scope>) -> LiftInSubquery<'scope> {
LiftInSubquery {
cx: self.cx.scope(target),
exclude_nulls: self.exclude_nulls,
}
}
fn exclude_nulls(&self, expr: &mut stmt::Expr) {
if !self.exclude_nulls {
return;
}
let stmt::Expr::InSubquery(expr) = expr else {
return;
};
let select = expr.query.body.as_select_mut_unwrap();
let target = select.source.model_id_unwrap();
let returning = select.returning.as_project_unwrap().clone();
for field in returning
.as_record()
.map_or(std::slice::from_ref(&returning), |record| &record.fields)
{
let stmt::Expr::Reference(stmt::ExprReference::Field { index, .. }) = field else {
unreachable!();
};
if self.cx.schema().app.field(target.field(*index)).nullable {
select.add_filter(stmt::Expr::is_not_null(field.clone()));
}
}
}
}
impl VisitMut for LiftInSubquery<'_> {
fn visit_expr_mut(&mut self, expr: &mut stmt::Expr) {
let lifted = match expr {
stmt::Expr::InSubquery(e) if !e.negated => {
lift_in_subquery(&self.cx, &e.expr, &e.query)
}
stmt::Expr::BinaryOp(_) | stmt::Expr::Like(_) | stmt::Expr::IsVariant(_) => {
try_lift_relation_path_predicate(&self.cx, expr)
}
_ => None,
};
if let Some(mut lifted) = lifted {
self.exclude_nulls(&mut lifted);
*expr = lifted;
}
stmt::visit_mut::visit_expr_mut(self, expr);
}
fn visit_stmt_delete_mut(&mut self, stmt: &mut stmt::Delete) {
self.visit_source_mut(&mut stmt.from);
let mut s = self.scope(&stmt.from);
s.visit_filter_mut(&mut stmt.filter);
if let Some(returning) = &mut stmt.returning {
s.visit_returning_mut(returning);
}
}
fn visit_stmt_insert_mut(&mut self, stmt: &mut stmt::Insert) {
self.visit_insert_target_mut(&mut stmt.target);
let mut s = self.scope(&stmt.target);
s.visit_stmt_query_mut(&mut stmt.source);
if let Some(returning) = &mut stmt.returning {
s.visit_returning_mut(returning);
}
}
fn visit_stmt_select_mut(&mut self, stmt: &mut stmt::Select) {
self.visit_source_mut(&mut stmt.source);
let mut s = self.scope(&stmt.source);
s.visit_filter_mut(&mut stmt.filter);
s.visit_returning_mut(&mut stmt.returning);
}
fn visit_stmt_update_mut(&mut self, stmt: &mut stmt::Update) {
self.visit_update_target_mut(&mut stmt.target);
let mut s = self.scope(&stmt.target);
s.visit_assignments_mut(&mut stmt.assignments);
s.visit_filter_mut(&mut stmt.filter);
if let Some(expr) = &mut stmt.condition.expr {
s.visit_expr_mut(expr);
}
if let Some(returning) = &mut stmt.returning {
s.visit_returning_mut(returning);
}
}
}
struct LiftBelongsTo<'a> {
cx: ExprContext<'a>,
belongs_to: &'a BelongsTo,
fk_field_matches: Vec<bool>,
fail: bool,
operands: Vec<stmt::Expr>,
}
pub(super) fn lift_in_subquery(
cx: &ExprContext,
expr: &stmt::Expr,
query: &stmt::Query,
) -> Option<stmt::Expr> {
match expr {
stmt::Expr::Project(_) | stmt::Expr::Variant(_) => {
lift_projection_in_subquery(cx, expr, query)
}
stmt::Expr::Reference(expr_reference @ stmt::ExprReference::Field { .. }) => {
let field = cx.resolve_expr_reference(expr_reference).as_field_unwrap();
lift_relation_in_subquery(cx, &Relation::Field(field), query)
}
_ => None,
}
}
struct RelationPath<'a> {
relation: Relation<'a>,
target: ModelId,
target_expr: Option<Expr>,
}
enum Relation<'a> {
Field(&'a app::Field),
Embedded {
belongs_to: &'a BelongsTo,
key_expr: Expr,
},
}
fn resolve_relation_path<'a>(cx: &ExprContext<'a>, expr: &Expr) -> Option<RelationPath<'a>> {
let resolved = relation_expr::resolve(cx, expr)?;
let relation = match resolved.embedded() {
Some(belongs_to) => Relation::Embedded {
belongs_to,
key_expr: resolved.key_expr()?,
},
None => Relation::Field(resolved.field),
};
Some(RelationPath {
relation,
target: resolved.target,
target_expr: resolved.target_expr(),
})
}
fn lift_relation_in_subquery(
cx: &ExprContext,
relation: &Relation<'_>,
query: &stmt::Query,
) -> Option<stmt::Expr> {
match relation {
Relation::Field(field) => match &field.ty {
FieldTy::BelongsTo(belongs_to) => lift_belongs_to_in_subquery(cx, belongs_to, query),
FieldTy::Has(has) => {
lift_has_n_in_subquery(has.target, has.pair(&cx.schema().app), query)
}
FieldTy::Via(via) => lift_via_in_subquery(cx, via, query),
_ => None,
},
Relation::Embedded {
belongs_to,
key_expr,
} => lift_fk_in_subquery(
belongs_to.target,
key_expr.clone(),
super::key_field_refs(0, belongs_to.foreign_key.fields.iter().map(|fk| fk.target)),
query,
),
}
}
fn lift_via_in_subquery(
cx: &ExprContext,
via: &app::Via,
query: &stmt::Query,
) -> Option<stmt::Expr> {
if via.is_scalar() {
return None;
}
let fields = super::relation_path::flatten_via_path(cx.schema(), via)?;
let target = fields
.last()
.and_then(|field| cx.schema().app.field(*field).relation_target_id())?;
if target != via.target || target != query.body.as_select_unwrap().source.model_id_unwrap() {
return None;
}
let (base, projection) = fields.split_first()?;
let base = stmt::Expr::ref_self_field(*base);
let path = if projection.is_empty() {
base
} else {
let projection = projection
.iter()
.map(|field| field.index)
.collect::<Vec<_>>();
stmt::Expr::project(base, projection.as_slice())
};
lift_in_subquery(cx, &path, query)
}
fn lift_projection_in_subquery(
cx: &ExprContext,
path: &stmt::Expr,
query: &stmt::Query,
) -> Option<stmt::Expr> {
let RelationPath {
relation,
target,
target_expr,
} = resolve_relation_path(cx, path)?;
let Some(inner_lhs) = target_expr else {
return lift_relation_in_subquery(cx, &relation, query);
};
if let Some(index) = inner_lhs.as_self_field_index()
&& let Relation::Field(field) = relation
&& let Some(direct) = try_fuse_paired_relations(cx, field, target, index, query)
{
return Some(direct);
}
let new_subquery = stmt::Query::new_select(
stmt::Source::from(target),
Expr::in_subquery(inner_lhs, query.clone()),
);
lift_relation_in_subquery(cx, &relation, &new_subquery)
}
fn try_fuse_paired_relations(
cx: &ExprContext,
outer_field: &app::Field,
target_model_id: ModelId,
head_idx: usize,
query: &stmt::Query,
) -> Option<stmt::Expr> {
let outer_belongs_to = match &outer_field.ty {
FieldTy::BelongsTo(rel) => rel,
_ => return None,
};
let target_model = cx.schema().app.model(target_model_id).as_root_unwrap();
let head_field = target_model.fields.get(head_idx)?;
let inner_has = match &head_field.ty {
FieldTy::Has(has) => has,
_ => return None,
};
let inner_pair = inner_has.pair(&cx.schema().app);
if outer_belongs_to.foreign_key.fields.len() != inner_pair.foreign_key.fields.len() {
return None;
}
for (outer_fk, inner_fk) in outer_belongs_to
.foreign_key
.fields
.iter()
.zip(inner_pair.foreign_key.fields.iter())
{
if outer_fk.target != inner_fk.target {
return None;
}
}
lift_fk_in_subquery(
inner_has.target,
super::key_field_refs(
0,
outer_belongs_to
.foreign_key
.fields
.iter()
.map(|fk| fk.source),
),
super::key_field_refs(0, inner_pair.foreign_key.fields.iter().map(|fk| fk.source)),
query,
)
}
pub(super) fn try_lift_relation_path_predicate(
cx: &ExprContext,
expr: &stmt::Expr,
) -> Option<stmt::Expr> {
let (path, filter) = resolve_relation_predicate(cx, expr)?;
lift_relation_predicate(cx, &path, filter)
}
fn resolve_relation_predicate<'a>(
cx: &ExprContext<'a>,
expr: &stmt::Expr,
) -> Option<(RelationPath<'a>, stmt::Expr)> {
match expr {
Expr::BinaryOp(e) => {
let comparison = |op: stmt::BinaryOp, subject: &Expr, other: &Expr| {
resolve_relation_path_predicate(cx, subject, |target_lhs| {
Expr::binary_op(target_lhs, op, other.clone())
})
};
comparison(e.op, &e.lhs, &e.rhs).or_else(|| comparison(e.op.commute()?, &e.rhs, &e.lhs))
}
Expr::Like(e) => resolve_relation_path_predicate(cx, &e.expr, |target_lhs| {
stmt::ExprLike {
expr: Box::new(target_lhs),
pattern: e.pattern.clone(),
escape: e.escape,
case_insensitive: e.case_insensitive,
}
.into()
}),
Expr::IsVariant(e) => resolve_relation_path_predicate(cx, &e.expr, |target_lhs| {
Expr::is_variant(target_lhs, e.variant)
}),
_ => None,
}
}
fn resolve_relation_path_predicate<'a>(
cx: &ExprContext<'a>,
subject: &stmt::Expr,
make_filter: impl FnOnce(stmt::Expr) -> stmt::Expr,
) -> Option<(RelationPath<'a>, stmt::Expr)> {
let path = resolve_relation_path(cx, subject)?;
let filter = make_filter(path.target_expr.clone()?);
Some((path, filter))
}
fn lift_relation_predicate(
cx: &ExprContext,
path: &RelationPath<'_>,
filter: stmt::Expr,
) -> Option<stmt::Expr> {
let subquery = stmt::Query::new_select(stmt::Source::from(path.target), filter);
lift_relation_in_subquery(cx, &path.relation, &subquery)
}
fn lift_fk_in_subquery(
target: ModelId,
lhs: stmt::Expr,
returning: stmt::Expr,
query: &stmt::Query,
) -> Option<stmt::Expr> {
if target != query.body.as_select_unwrap().source.model_id_unwrap() {
return None;
}
let mut subquery = query.clone();
subquery.body.as_select_mut_unwrap().returning = stmt::Returning::Project(returning);
Some(stmt::Expr::in_subquery(lhs, subquery))
}
fn lift_belongs_to_in_subquery(
cx: &ExprContext,
belongs_to: &BelongsTo,
query: &stmt::Query,
) -> Option<stmt::Expr> {
if belongs_to.target != query.body.as_select_unwrap().source.model_id_unwrap() {
return None;
}
let select = query.body.as_select_unwrap();
let mut lift = LiftBelongsTo {
cx: cx.scope(&select.source),
belongs_to,
fk_field_matches: vec![false; belongs_to.foreign_key.fields.len()],
operands: vec![],
fail: false,
};
lift.visit_filter(&select.filter);
let all_fks_matched = lift.fk_field_matches.iter().all(|m| *m);
if lift.fail || !all_fks_matched {
lift_fk_in_subquery(
belongs_to.target,
super::key_field_refs(0, belongs_to.foreign_key.fields.iter().map(|fk| fk.source)),
super::key_field_refs(0, belongs_to.foreign_key.fields.iter().map(|fk| fk.target)),
query,
)
} else {
Some(if lift.operands.len() == 1 {
lift.operands.into_iter().next().unwrap()
} else {
stmt::ExprAnd {
operands: lift.operands,
}
.into()
})
}
}
fn lift_has_n_in_subquery(
target: ModelId,
pair: &BelongsTo,
query: &stmt::Query,
) -> Option<stmt::Expr> {
lift_fk_in_subquery(
target,
super::key_field_refs(0, pair.foreign_key.fields.iter().map(|fk| fk.target)),
super::key_field_refs(0, pair.foreign_key.fields.iter().map(|fk| fk.source)),
query,
)
}
impl Visit for LiftBelongsTo<'_> {
fn visit_expr_in_subquery(&mut self, _i: &stmt::ExprInSubquery) {
}
fn visit_expr_binary_op(&mut self, i: &stmt::ExprBinaryOp) {
match (&*i.lhs, &*i.rhs) {
(stmt::Expr::Reference(expr_reference), other)
| (other, stmt::Expr::Reference(expr_reference)) => {
assert!(i.op.is_eq() || i.op.is_ne());
if i.op.is_eq() || i.op.is_ne() {
let field = self
.cx
.resolve_expr_reference(expr_reference)
.as_field_unwrap();
self.lift_fk_constraint(field.id, i.op, other);
} else {
self.fail = true;
}
}
_ => {
self.fail = true;
}
}
}
}
impl LiftBelongsTo<'_> {
fn lift_fk_constraint(&mut self, field: FieldId, op: stmt::BinaryOp, expr: &stmt::Expr) {
for (i, fk_field) in self.belongs_to.foreign_key.fields.iter().enumerate() {
if fk_field.target == field {
if self.fk_field_matches[i] {
todo!("not handled");
}
self.operands.push(stmt::Expr::binary_op(
stmt::Expr::ref_self_field(fk_field.source),
op,
expr.clone(),
));
self.fk_field_matches[i] = true;
return;
}
}
self.fail = true;
}
}