use crate::DjogiError;
use crate::context::ContextInner;
use crate::model::Model;
use crate::pg::accumulator::SqlAccumulator;
use crate::pg::decode::FromJoinedPgRow;
use crate::relation::joined_row::JoinedRow;
use std::any::Any;
use std::collections::HashMap;
use tokio_postgres::Row as PgRow;
pub(crate) type JoinDecoderFn =
for<'r> fn(
row: &'r PgRow,
prefix: &str,
) -> Result<Option<Box<dyn Any + Send + Sync>>, DjogiError>;
pub(crate) type ChildDescriptorFn = fn() -> &'static crate::descriptor::ModelDescriptor;
#[derive(Clone)]
pub(crate) struct ErasedSelectRelated {
pub source_column: &'static str,
pub child_table: &'static str,
pub decoder: JoinDecoderFn,
pub child_descriptor: ChildDescriptorFn,
}
pub(crate) fn child_descriptor<Child: Model>() -> &'static crate::descriptor::ModelDescriptor {
<Child as Model>::descriptor()
}
impl std::fmt::Debug for ErasedSelectRelated {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ErasedSelectRelated")
.field("source_column", &self.source_column)
.field("child_table", &self.child_table)
.finish_non_exhaustive()
}
}
pub(crate) fn join_decoder<Child>(
row: &PgRow,
prefix: &str,
) -> Result<Option<Box<dyn Any + Send + Sync>>, DjogiError>
where
Child: Model + FromJoinedPgRow + Send + Sync + 'static,
{
let id_alias = format!("{prefix}id");
let id_probe: Result<Option<i64>, _> = row.try_get(id_alias.as_str());
debug_assert!(
id_probe.is_ok(),
"select_related: missing join alias '{id_alias}' on joined row — framework bug (check select_columns emission vs push_joins alias)"
);
let child_is_null = id_probe.map(|v| v.is_none()).unwrap_or(true);
if child_is_null {
return Ok(None);
}
let child = <Child as FromJoinedPgRow>::from_joined_pg_row(row, prefix)?;
Ok(Some(Box::new(child) as Box<dyn Any + Send + Sync>))
}
pub(crate) fn push_joins<T: Model>(acc: &mut SqlAccumulator, paths: &[ErasedSelectRelated]) {
for path in paths {
acc.push_sql(" LEFT JOIN ");
acc.push_sql(path.child_table);
acc.push_sql(" rel_");
acc.push_sql(path.source_column);
acc.push_sql(" ON ");
acc.push_sql(T::table_name());
acc.push_sql(".");
acc.push_sql(path.source_column);
acc.push_sql(" = rel_");
acc.push_sql(path.source_column);
acc.push_sql(".id");
}
}
pub(crate) fn select_columns<T: Model>(paths: &[ErasedSelectRelated]) -> String {
let mut out = String::new();
out.push_str(T::table_name());
out.push_str(".*");
for path in paths {
let alias = format!("rel_{}", path.source_column);
let desc = (path.child_descriptor)();
for field in desc.fields {
crate::ident::debug_assert_ident!(field.name, "field_name");
out.push_str(", ");
out.push_str(&alias);
out.push('.');
out.push_str(field.name);
out.push_str(" AS \"");
out.push_str(&alias);
out.push('.');
out.push_str(field.name);
out.push('"');
}
}
out
}
pub(crate) fn decode_joined_row<T: Model + FromJoinedPgRow>(
row: &PgRow,
paths: &[ErasedSelectRelated],
) -> Result<JoinedRow<T>, DjogiError> {
let parent = <T as FromJoinedPgRow>::from_joined_pg_row(row, "")?;
let mut relations: HashMap<&'static str, Box<dyn Any + Send + Sync>> =
HashMap::with_capacity(paths.len());
for path in paths {
let prefix = format!("rel_{}.", path.source_column);
if let Some(child_box) = (path.decoder)(row, &prefix)? {
relations.insert(path.source_column, child_box);
}
}
Ok(JoinedRow::new(parent, relations))
}
pub(crate) fn apply_select_related<T>(
rows: Vec<PgRow>,
paths: &[ErasedSelectRelated],
) -> Result<Vec<JoinedRow<T>>, DjogiError>
where
T: Model + FromJoinedPgRow,
{
let mut out: Vec<JoinedRow<T>> = Vec::with_capacity(rows.len());
for row in &rows {
out.push(decode_joined_row::<T>(row, paths)?);
}
Ok(out)
}
pub(crate) async fn stitch_prefetches_into_joined<T>(
mut joined: Vec<JoinedRow<T>>,
prefetches: &[crate::relation::prefetch::ErasedPrefetch],
exec: &mut ContextInner,
) -> Result<Vec<JoinedRow<T>>, DjogiError>
where
T: Model,
T::Pk: Clone + Send + Sync + 'static,
{
if prefetches.is_empty() || joined.is_empty() {
return Ok(joined);
}
for prefetch in prefetches {
let parent_pks_for_loader: Vec<Box<dyn Any + Send + Sync>> = joined
.iter()
.map(|jr| Box::new(jr.row.pk_value().clone()) as Box<dyn Any + Send + Sync>)
.collect();
let aligned = (prefetch.loader)(
&mut *exec,
prefetch.parent_table,
prefetch.source_column,
parent_pks_for_loader,
)
.await?;
for (jr, slot) in joined.iter_mut().zip(aligned) {
if let Some(child_box) = slot {
jr.relations_mut().insert(prefetch.source_column, child_box);
}
}
}
Ok(joined)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::descriptor::{
FieldDescriptor, FieldSqlType, ModelDescriptor, PkType, field_descriptor, model_descriptor,
};
use crate::model::Model;
use crate::pg::accumulator::SqlAccumulator;
use crate::types::HeerId;
use std::future::Future;
struct Src;
impl crate::model::__sealed::Sealed for Src {}
#[allow(clippy::manual_async_fn)]
impl Model for Src {
type Pk = HeerId;
type Fields = ();
fn table_name() -> &'static str {
"srcs"
}
fn pk_value(&self) -> &HeerId {
unreachable!()
}
fn descriptor() -> &'static ModelDescriptor {
unreachable!()
}
fn get(
_ctx: &mut crate::context::DjogiContext,
_id: HeerId,
) -> impl Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn create(
_ctx: &mut crate::context::DjogiContext,
_v: Self,
) -> impl Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn save<'ctx>(
&'ctx mut self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl Future<Output = Result<(), crate::DjogiError>> + Send + 'ctx {
async { unreachable!() }
}
fn delete(
self,
_ctx: &mut crate::context::DjogiContext,
) -> impl Future<Output = Result<(), crate::DjogiError>> + Send {
async { unreachable!() }
}
fn refresh_from_db<'ctx>(
&'ctx self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl Future<Output = Result<Self, crate::DjogiError>> + Send + 'ctx {
async { unreachable!() }
}
}
fn dummy_decoder(
_row: &PgRow,
_prefix: &str,
) -> Result<Option<Box<dyn Any + Send + Sync>>, DjogiError> {
unreachable!("dummy decoder should not run in SQL-emission tests")
}
static OWNERS_DESC: ModelDescriptor = ModelDescriptor {
..model_descriptor(
"Owner",
"owners",
PkType::HeerId,
&[
FieldDescriptor {
unique: true,
indexed: true,
..field_descriptor("id", FieldSqlType::BigInt, false)
},
FieldDescriptor {
..field_descriptor("name", FieldSqlType::Text, false)
},
],
)
};
static FUEL_TYPES_DESC: ModelDescriptor = ModelDescriptor {
..model_descriptor(
"FuelType",
"fuel_types",
PkType::HeerId,
&[FieldDescriptor {
unique: true,
indexed: true,
..field_descriptor("id", FieldSqlType::BigInt, false)
}],
)
};
fn owners_descriptor() -> &'static ModelDescriptor {
&OWNERS_DESC
}
fn fuel_types_descriptor() -> &'static ModelDescriptor {
&FUEL_TYPES_DESC
}
#[test]
fn push_joins_emits_left_join_with_aliased_table() {
let path = ErasedSelectRelated {
source_column: "owner_id",
child_table: "owners",
decoder: dummy_decoder,
child_descriptor: owners_descriptor,
};
let mut acc = SqlAccumulator::new("SELECT * FROM srcs");
push_joins::<Src>(&mut acc, &[path]);
let sql = acc.sql();
assert!(
sql.contains("LEFT JOIN owners rel_owner_id ON srcs.owner_id = rel_owner_id.id"),
"expected aliased LEFT JOIN, got: {sql}"
);
}
#[test]
fn push_joins_emits_one_clause_per_path() {
let paths = vec![
ErasedSelectRelated {
source_column: "owner_id",
child_table: "owners",
decoder: dummy_decoder,
child_descriptor: owners_descriptor,
},
ErasedSelectRelated {
source_column: "fuel_type_id",
child_table: "fuel_types",
decoder: dummy_decoder,
child_descriptor: fuel_types_descriptor,
},
];
let mut acc = SqlAccumulator::new("SELECT * FROM srcs");
push_joins::<Src>(&mut acc, &paths);
let sql = acc.sql();
assert!(
sql.contains("LEFT JOIN owners rel_owner_id"),
"missing owner join in: {sql}"
);
assert!(
sql.contains("LEFT JOIN fuel_types rel_fuel_type_id"),
"missing fuel_type join in: {sql}"
);
}
#[test]
fn select_columns_emits_parent_star_and_aliased_children() {
let path = ErasedSelectRelated {
source_column: "owner_id",
child_table: "owners",
decoder: dummy_decoder,
child_descriptor: owners_descriptor,
};
let cols = select_columns::<Src>(&[path]);
assert!(cols.starts_with("srcs.*"), "got: {cols}");
assert!(
cols.contains("rel_owner_id.id AS \"rel_owner_id.id\""),
"missing aliased id column, got: {cols}"
);
assert!(
cols.contains("rel_owner_id.name AS \"rel_owner_id.name\""),
"missing aliased name column, got: {cols}"
);
}
#[test]
fn select_columns_empty_paths_returns_just_parent_star() {
let cols = select_columns::<Src>(&[]);
assert_eq!(cols, "srcs.*");
}
}