use toasty_core::{
schema::{app, mapping},
stmt,
};
use crate::engine::lower::LowerStatement;
use crate::schema::lazy_slot;
struct FlatInclude {
projection: stmt::Projection,
query: Option<stmt::Query>,
}
#[derive(Default)]
struct IncludeQuery {
filter: Option<stmt::Expr>,
order_by: Option<stmt::OrderBy>,
}
struct FieldIncludes {
include_self: bool,
top_query: IncludeQuery,
sub_paths: Vec<FlatInclude>,
}
impl LowerStatement<'_, '_> {
pub(super) fn process_top_level_includes(
&mut self,
returning: &mut stmt::Expr,
includes: &[stmt::Include],
is_insert: bool,
) {
let stmt::Expr::Record(record) = returning else {
return;
};
let flat: Vec<FlatInclude> = includes.iter().map(flatten_include).collect();
let app_fields = &self.model_unwrap().fields;
let mapping_fields = &self.mapping_unwrap().fields;
self.process_fields(record, app_fields, mapping_fields, &flat, is_insert);
}
fn process_fields(
&mut self,
returning: &mut stmt::ExprRecord,
app_fields: &[app::Field],
mapping_fields: &[mapping::Field],
includes: &[FlatInclude],
is_insert: bool,
) {
for (i, (field, mapping)) in app_fields.iter().zip(mapping_fields).enumerate() {
let field_includes = partition_includes(includes, i);
if field.ty.is_relation() {
if field_includes.self_included() {
self.build_include_subquery(
returning,
i,
&field_includes.sub_paths,
field_includes.top_query,
);
}
continue;
}
self.process_field(
&mut returning[i],
field,
mapping,
&field_includes,
is_insert,
);
}
}
fn process_field(
&mut self,
returning: &mut stmt::Expr,
field: &app::Field,
mapping: &mapping::Field,
matches: &FieldIncludes,
is_insert: bool,
) {
if field.deferred {
if !is_insert && !matches.self_included() {
return;
}
*returning = lazy_slot::loaded_expr(loaded_form(field, mapping));
if let app::FieldTy::Embedded(embedded) = &field.ty {
let stmt::Expr::Record(outer) = returning else {
unreachable!("just-wrapped record");
};
self.process_embed(
&mut outer[0],
embedded.target,
mapping,
&matches.sub_paths,
is_insert,
);
}
return;
}
if let app::FieldTy::Embedded(embedded) = &field.ty
&& (is_insert || !matches.sub_paths.is_empty())
{
self.process_embed(
returning,
embedded.target,
mapping,
&matches.sub_paths,
is_insert,
);
}
}
fn process_embed(
&mut self,
returning: &mut stmt::Expr,
target: app::ModelId,
mapping: &mapping::Field,
sub_includes: &[FlatInclude],
is_insert: bool,
) {
match (self.schema().app.model(target), mapping) {
(app::Model::EmbeddedStruct(em), mapping::Field::Struct(fs)) => {
let record = match returning {
stmt::Expr::Record(record) => record,
stmt::Expr::Match(match_expr) => {
let Some(stmt::Expr::Record(record)) =
match_expr.arms.first_mut().map(|arm| &mut arm.expr)
else {
return;
};
record
}
_ => return,
};
self.process_fields(
record,
em.fields.as_slice(),
fs.fields.as_slice(),
sub_includes,
is_insert,
);
}
(app::Model::EmbeddedEnum(em), mapping::Field::Enum(fe)) => {
self.process_enum_arms(returning, em, fe, sub_includes, is_insert);
}
_ => {}
}
}
fn process_enum_arms(
&mut self,
returning: &mut stmt::Expr,
app_enum: &app::EmbeddedEnum,
mapping: &mapping::FieldEnum,
sub_includes: &[FlatInclude],
is_insert: bool,
) {
let stmt::Expr::Match(match_expr) = returning else {
return;
};
for (variant_idx, arm) in match_expr.arms.iter_mut().enumerate() {
let variant_fields: Vec<&app::Field> = app_enum.variant_fields(variant_idx).collect();
if variant_fields.is_empty() {
continue;
}
let stmt::Expr::Record(arm_record) = &mut arm.expr else {
continue;
};
let variant_mapping = &mapping.variants[variant_idx];
let arm_sub_includes: Vec<FlatInclude> = sub_includes
.iter()
.filter_map(|fi| {
let (first, rest) = fi.projection.as_slice().split_first()?;
(*first == variant_idx).then(|| FlatInclude {
projection: stmt::Projection::from(rest),
query: fi.query.clone(),
})
})
.collect();
for (j, (var_field, var_mapping)) in variant_fields
.iter()
.zip(&variant_mapping.fields)
.enumerate()
{
let field_includes = partition_includes(&arm_sub_includes, j);
self.process_field(
&mut arm_record[j + 1],
var_field,
var_mapping,
&field_includes,
is_insert,
);
}
}
}
fn build_include_subquery(
&mut self,
returning: &mut stmt::ExprRecord,
field_index: usize,
nested: &[FlatInclude],
top_query: IncludeQuery,
) {
let value = self.build_relation_subquery_inner(field_index, nested, top_query);
returning[field_index] = if self.model_unwrap().fields[field_index].deferred {
lazy_slot::loaded_expr(value)
} else {
value
};
}
pub(super) fn build_relation_subquery(&mut self, field_index: usize) -> stmt::Expr {
self.build_relation_subquery_inner(field_index, &[], IncludeQuery::default())
}
fn build_relation_subquery_inner(
&mut self,
field_index: usize,
nested: &[FlatInclude],
top_query: IncludeQuery,
) -> stmt::Expr {
let field = &self.model_unwrap().fields[field_index];
let via = match &field.ty {
app::FieldTy::Via(via) => Some(via),
_ => None,
};
if let Some(via) = via {
if !self.capability().sql {
todo!(
"`.include()` / `.select()` of a multi-step `via` relation is only \
supported on SQL backends; query the relation directly instead"
);
}
if top_query.filter.is_some()
|| top_query.order_by.is_some()
|| nested.iter().any(|fi| query_has_modifiers(&fi.query))
{
todo!(
"include query modifiers on a multi-step `via` relation are not yet supported"
);
}
let nested_projections: Vec<stmt::Projection> =
nested.iter().map(|fi| fi.projection.clone()).collect();
return self.build_via_include_subquery(field_index, via, &nested_projections);
}
let (mut stmt, target_model_id) = match &field.ty {
app::FieldTy::Has(rel) => {
let mut query = stmt::Query::new_select(
rel.target,
stmt::Expr::eq(
stmt::Expr::ref_parent_model(),
stmt::Expr::ref_self_field(rel.pair_id),
),
);
if rel.is_one() {
query.single = true;
}
(query, rel.target)
}
app::FieldTy::BelongsTo(rel) => {
let source_fk;
let target_pk;
if let [fk_field] = &rel.foreign_key.fields[..] {
source_fk = stmt::Expr::ref_parent_field(fk_field.source);
target_pk = stmt::Expr::ref_self_field(fk_field.target);
} else {
let mut source_fk_fields = vec![];
let mut target_pk_fields = vec![];
for fk_field in &rel.foreign_key.fields {
source_fk_fields.push(stmt::Expr::ref_parent_field(fk_field.source));
target_pk_fields.push(stmt::Expr::ref_self_field(fk_field.target));
}
source_fk = stmt::Expr::record_from_vec(source_fk_fields);
target_pk = stmt::Expr::record_from_vec(target_pk_fields);
}
let mut query =
stmt::Query::new_select(rel.target, stmt::Expr::eq(source_fk, target_pk));
query.single = true;
(query, rel.target)
}
_ => unreachable!("build_include_subquery called on non-relation field"),
};
if let Some(filter) = top_query.filter {
stmt.add_filter(filter);
}
stmt.order_by = top_query.order_by;
for fi in nested {
if !fi.projection.is_empty() {
stmt.include(stmt::Include {
path: stmt::Path {
root: stmt::PathRoot::Model(target_model_id),
projection: fi.projection.clone(),
},
query: fi.query.clone(),
});
}
}
self.lower_sub_stmt(stmt::Statement::Query(stmt))
}
}
impl FieldIncludes {
fn self_included(&self) -> bool {
self.include_self || !self.sub_paths.is_empty()
}
}
fn partition_includes(includes: &[FlatInclude], i: usize) -> FieldIncludes {
let mut include_self = false;
let mut unfiltered_self = false;
let mut top_filter: Option<stmt::Expr> = None;
let mut top_order_by = None;
let mut sub_paths = Vec::new();
for fi in includes {
if let Some((first, rest)) = fi.projection.as_slice().split_first()
&& *first == i
{
if rest.is_empty() {
include_self = true;
top_order_by = fi.query.as_ref().and_then(|query| query.order_by.clone());
match query_filter_expr(&fi.query) {
Some(f) if !unfiltered_self => {
top_filter = Some(match top_filter.take() {
Some(prev) => stmt::Expr::or(prev, f),
None => f,
});
}
Some(_) => {}
None => {
unfiltered_self = true;
top_filter = None;
}
}
} else {
sub_paths.push(FlatInclude {
projection: stmt::Projection::from(rest),
query: fi.query.clone(),
});
}
}
}
FieldIncludes {
include_self,
top_query: IncludeQuery {
filter: top_filter,
order_by: top_order_by,
},
sub_paths,
}
}
fn flatten_include(include: &stmt::Include) -> FlatInclude {
FlatInclude {
projection: flatten_path(&include.path),
query: include.query.clone(),
}
}
fn query_filter_expr(query: &Option<stmt::Query>) -> Option<stmt::Expr> {
match &query.as_ref()?.body {
stmt::ExprSet::Select(select) => select.filter.expr.clone(),
_ => None,
}
}
fn query_has_modifiers(query: &Option<stmt::Query>) -> bool {
query.as_ref().is_some_and(|query| {
query_filter_expr(&Some(query.clone())).is_some() || query.order_by.is_some()
})
}
fn flatten_path(path: &stmt::Path) -> stmt::Projection {
let mut acc = stmt::Projection::identity();
push_root_steps(&path.root, &mut acc);
for step in path.projection.as_slice() {
acc.push(*step);
}
acc
}
fn push_root_steps(root: &stmt::PathRoot, acc: &mut stmt::Projection) {
if let stmt::PathRoot::Variant { parent, variant_id } = root {
push_root_steps(&parent.root, acc);
for step in parent.projection.as_slice() {
acc.push(*step);
}
acc.push(variant_id.index);
}
}
fn loaded_form(field: &app::Field, mapping: &mapping::Field) -> stmt::Expr {
match (&field.ty, mapping) {
(app::FieldTy::Primitive(_), mapping::Field::Primitive(p)) => p.column_expr.clone(),
(app::FieldTy::Embedded(_), mapping::Field::Struct(s)) => s.default_returning.clone(),
(app::FieldTy::Embedded(_), mapping::Field::Enum(e)) => e.default_returning.clone(),
_ => unreachable!("deferred field has unexpected mapping shape"),
}
}