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 {
included: bool,
top_query: IncludeQuery,
sub_paths: Vec<FlatInclude>,
}
impl LowerStatement<'_, '_> {
pub(super) fn process_projected_returning(&mut self, expr: &mut stmt::Expr) {
if self.process_projected_field(expr) {
return;
}
struct ProjectionVisitor<'a, 'b, 'c>(&'a mut LowerStatement<'b, 'c>);
impl stmt::VisitMut for ProjectionVisitor<'_, '_, '_> {
fn visit_expr_mut(&mut self, expr: &mut stmt::Expr) {
self.0.process_projected_returning(expr);
}
fn visit_stmt_mut(&mut self, _: &mut stmt::Statement) {}
fn visit_stmt_query_mut(&mut self, _: &mut stmt::Query) {}
}
stmt::visit_mut::visit_expr_mut(&mut ProjectionVisitor(self), expr);
}
fn process_projected_field(&mut self, expr: &mut stmt::Expr) -> bool {
let Some(projected) = self.expr_cx.resolve_projected_field(expr) else {
return false;
};
if projected.nesting != 0 {
return false;
}
match &projected.field.ty {
app::FieldTy::Embedded(embedded) => {
use stmt::VisitMut;
self.visit_expr_mut(expr);
super::Simplify::with_context(self.expr_cx, self.capability()).visit_expr_mut(expr);
self.process_embed(
expr,
embedded.target,
projected.mapping,
&[],
&projected.path,
);
}
_ if projected.field.ty.is_relation() && expr.is_self_field() => {
*expr = self.build_relation_subquery(projected.field.id.index);
}
_ => {}
}
true
}
pub(super) fn process_update_embedded_relations(&mut self, returning: &mut stmt::Returning) {
let stmt::Returning::Project(expr) = returning else {
return;
};
let Some(model) = self.model() else {
return;
};
let mapping = self.mapping_unwrap();
self.lower_returning().process_sparse_embeds_for_update(
expr,
&model.fields,
&mapping.fields,
&stmt::Path::model(model.id),
);
}
fn process_sparse_embeds_for_update(
&mut self,
expr: &mut stmt::Expr,
app_fields: &[app::Field],
mapping_fields: &[mapping::Field],
record_path: &stmt::Path,
) {
let stmt::Expr::Cast(cast) = expr else {
return;
};
let stmt::Type::SparseRecord(returned_fields) = &mut cast.ty else {
return;
};
let stmt::Expr::Record(record) = &mut *cast.expr else {
return;
};
if !record_path.projection.is_empty() {
self.refresh_changed_relations(
returned_fields,
&mut record.fields,
app_fields,
record_path,
);
}
for (field_index, field_expr) in returned_fields.iter().zip(&mut record.fields) {
let app::FieldTy::Embedded(embedded) = &app_fields[field_index].ty else {
continue;
};
let embed_path = field_path(record_path, field_index);
if matches!(
field_expr,
stmt::Expr::Cast(cast) if matches!(cast.ty, stmt::Type::SparseRecord(_))
) {
let app::Model::EmbeddedStruct(model) = self.schema().app.model(embedded.target)
else {
unreachable!("only structs support partial updates")
};
self.process_sparse_embeds_for_update(
field_expr,
&model.fields,
&mapping_fields[field_index].as_struct().unwrap().fields,
&embed_path,
);
} else {
use stmt::VisitMut;
self.visit_expr_mut(field_expr);
self.process_embed(
field_expr,
embedded.target,
&mapping_fields[field_index],
&[],
&embed_path,
);
}
}
}
fn refresh_changed_relations(
&mut self,
returned_fields: &mut stmt::PathFieldSet,
returned_values: &mut Vec<stmt::Expr>,
app_fields: &[app::Field],
record_path: &stmt::Path,
) {
for (field_index, field) in app_fields.iter().enumerate() {
let app::FieldTy::BelongsTo(relation) = &field.ty else {
continue;
};
if field.deferred {
continue;
}
let has_returned_foreign_key_field = relation
.foreign_key
.fields
.iter()
.any(|fk| returned_fields.contains(fk.source.index));
if !has_returned_foreign_key_field {
continue;
}
let relation_load = self.build_relation_subquery_inner(
field,
record_path,
&[],
IncludeQuery::default(),
);
let value_index = returned_fields
.iter()
.take_while(|index| *index < field_index)
.count();
if returned_fields.contains(field_index) {
returned_values[value_index] = relation_load;
} else {
returned_fields.insert(field_index);
returned_values.insert(value_index, relation_load);
}
}
}
pub(super) fn process_top_level_includes(
&mut self,
record: &mut stmt::ExprRecord,
includes: &[stmt::Include],
) {
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(
&mut record.fields,
app_fields,
mapping_fields,
&flat,
&stmt::Path::model(self.model_unwrap().id),
);
}
fn process_fields(
&mut self,
returning: &mut [stmt::Expr],
app_fields: &[app::Field],
mapping_fields: &[mapping::Field],
includes: &[FlatInclude],
host: &stmt::Path,
) {
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 self.cx.is_insert_without_row() {
continue;
}
if self.cx.is_insert_with_row() && !field.ty.is_belongs_to() {
continue;
}
if field_includes.included || !field.deferred {
let value = self.build_relation_subquery_inner(
field,
host,
&field_includes.sub_paths,
field_includes.top_query,
);
returning[i] = if field.deferred {
lazy_slot::loaded_expr(value)
} else {
value
};
}
continue;
}
self.process_field(
&mut returning[i],
field,
mapping,
&field_includes,
&field_path(host, i),
);
}
}
fn process_field(
&mut self,
returning: &mut stmt::Expr,
field: &app::Field,
mapping: &mapping::Field,
matches: &FieldIncludes,
path: &stmt::Path,
) {
if field.deferred {
if !self.cx.is_insert() && !matches.included {
return;
}
if !self.cx.is_insert_with_row() {
*returning = lazy_slot::loaded_expr(loaded_form(field, mapping));
}
}
if let app::FieldTy::Embedded(embedded) = &field.ty {
let returning = if field.deferred {
let stmt::Expr::Record(outer) = returning else {
unreachable!("just-wrapped record");
};
&mut outer[0]
} else {
returning
};
self.process_embed(
returning,
embedded.target,
mapping,
&matches.sub_paths,
path,
);
}
}
fn process_embed(
&mut self,
returning: &mut stmt::Expr,
target: app::ModelId,
mapping: &mapping::Field,
sub_includes: &[FlatInclude],
path: &stmt::Path,
) {
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(
&mut record.fields,
em.fields.as_slice(),
fs.fields.as_slice(),
sub_includes,
path,
);
}
(app::Model::EmbeddedEnum(em), mapping::Field::Enum(fe)) => {
self.process_enum_arms(returning, em, fe, sub_includes, path);
}
_ => {}
}
}
fn process_enum_arms(
&mut self,
returning: &mut stmt::Expr,
app_enum: &app::EmbeddedEnum,
mapping: &mapping::FieldEnum,
sub_includes: &[FlatInclude],
path: &stmt::Path,
) {
let stmt::Expr::Match(match_expr) = returning else {
return;
};
for (variant_idx, arm) in match_expr.arms.iter_mut().enumerate() {
let variant_fields = app_enum.variant_fields(variant_idx);
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 host = stmt::Path::from_variant(
path.clone(),
app::VariantId {
model: app_enum.id,
index: variant_idx,
},
);
let matches = partition_includes(sub_includes, variant_idx);
self.process_fields(
&mut arm_record.fields[1..],
variant_fields,
&variant_mapping.fields,
&matches.sub_paths,
&host,
);
}
}
pub(super) fn build_relation_subquery(&mut self, field_index: usize) -> stmt::Expr {
self.build_relation_subquery_inner(
&self.model_unwrap().fields[field_index],
&stmt::Path::model(self.model_unwrap().id),
&[],
IncludeQuery::default(),
)
}
fn build_relation_subquery_inner(
&mut self,
field: &app::Field,
host: &stmt::Path,
nested: &[FlatInclude],
top_query: IncludeQuery,
) -> stmt::Expr {
let field_index = field.id.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 = super::scalar_or_record(
rel.foreign_key
.fields
.iter()
.map(|fk| self.relation_source_field(host, fk.source)),
);
let target_pk =
super::key_field_refs(0, rel.foreign_key.fields.iter().map(|fk| fk.target));
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(),
});
}
}
let mut statement = stmt::Statement::Query(stmt);
self.state
.engine
.normalize_stmt(&mut statement)
.expect("valid include subquery");
let relation_load = self.lower_sub_stmt(statement);
if self.cx.is_insert_with_row() && field.ty.is_belongs_to() {
self.order_relation_load_after_enclosing_inserts(&relation_load);
Self::single_relation_from_load(relation_load)
} else {
relation_load
}
}
fn relation_source_field(
&self,
record_path: &stmt::Path,
source_field: app::FieldId,
) -> stmt::Expr {
let local_index = match &record_path.root {
stmt::PathRoot::Variant { variant_id, .. } if record_path.projection.is_empty() => {
let app::Model::EmbeddedEnum(model) = self.schema().app.model(variant_id.model)
else {
unreachable!()
};
model
.variant_fields(variant_id.index)
.iter()
.position(|field| field.id == source_field)
.unwrap()
}
_ => source_field.index,
};
field_path(record_path, local_index).into_stmt_with_nesting(1)
}
}
fn field_path(host: &stmt::Path, index: usize) -> stmt::Path {
let mut path = host.clone();
path.projection.push(index);
path
}
fn partition_includes(includes: &[FlatInclude], i: usize) -> FieldIncludes {
let mut included = 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
{
included = true;
if rest.is_empty() {
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 => {
let f = f.clone();
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 {
included,
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.as_ref(),
_ => None,
}
}
fn query_has_modifiers(query: &Option<stmt::Query>) -> bool {
query_filter_expr(query).is_some()
|| query.as_ref().is_some_and(|query| query.order_by.is_some())
}
fn flatten_path(path: &stmt::Path) -> stmt::Projection {
let mut acc = if let stmt::PathRoot::Variant { parent, variant_id } = &path.root {
let mut acc = flatten_path(parent);
acc.push(variant_id.index);
acc
} else {
stmt::Projection::identity()
};
for step in path.projection.as_slice() {
acc.push(*step);
}
acc
}
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"),
}
}