#![allow(clippy::manual_async_fn)]
use crate::DjogiError;
use crate::context::DjogiContext;
use crate::model::Model;
use crate::pg::accumulator::{SqlAccumulator, as_params};
use crate::pg::decode::{FromJoinedPgRow, FromPgRow, try_get_scalar};
use crate::query::queryset::QuerySet;
use crate::query::sql::{build_count, build_exists, build_select, build_select_joined};
use crate::query::stream::{DEFAULT_FETCH_SIZE, ModelCursorStream, build_model_stream};
use crate::relation::joined_row::JoinedRow;
use crate::relation::prefetch::{PrefetchedRow, apply_prefetches};
use crate::relation::select_related::{apply_select_related, stitch_prefetches_into_joined};
use std::collections::HashMap;
use std::future::Future;
use std::hash::Hash;
pub(crate) async fn auto_set_tenant<T: Model>(ctx: &mut DjogiContext) -> Result<(), DjogiError> {
if T::descriptor().tenant_key.is_none() {
return Ok(());
}
if ctx.auth().is_none() {
return Ok(());
}
let tid: Option<String> = ctx.auth().and_then(|a| a.tenant_id.clone());
match tid {
Some(tid) => ctx.ensure_tenant_set(&tid).await?,
None => {
if ctx.applied_tenant_id().is_some() {
ctx.clear_tenant().await?;
}
if !ctx.__tenant_scope_suppressed_for_macros() {
tracing::warn!(
model = std::any::type_name::<T>(),
"auth attached but tenant_id is None on a tenant-keyed model; \
queries will span tenants — call ctx.with_no_tenant_scope() to suppress",
);
}
}
}
Ok(())
}
impl<T: Model> QuerySet<T>
where
T: FromPgRow,
{
#[doc(hidden)]
pub fn __sql_for_test(&self) -> Result<String, crate::DjogiError> {
let acc = crate::query::sql::build_select(self).map_err(crate::DjogiError::from)?;
let (sql, _binds) = acc.into_parts();
Ok(sql)
}
}
impl<T: Model> QuerySet<T>
where
T: FromPgRow + Send + Unpin,
{
pub fn fetch_all<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<Vec<T>, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
if self.is_empty() {
return Ok(Vec::new());
}
auto_set_tenant::<T>(ctx).await?;
let cache_target = self.cache_target.clone();
let acc = build_select(&self).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows = ctx.query_all(&sql, ¶ms).await?;
let result: Vec<T> = rows
.iter()
.map(|r| T::from_pg_row(r))
.collect::<Result<Vec<T>, _>>()?;
if let Some(target) = cache_target.as_ref() {
for row in &result {
target.insert(row).await;
}
}
Ok(result)
}
}
pub fn fetch_one<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<T, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
if self.is_empty() {
return Err(DjogiError::not_found(T::table_name()));
}
auto_set_tenant::<T>(ctx).await?;
let mut qs = self;
qs.limit = Some(2);
let cache_target = qs.cache_target.clone();
let acc = build_select(&qs).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows = ctx.query_all(&sql, ¶ms).await?;
match rows.len() {
0 => Err(DjogiError::not_found(T::table_name())),
1 => {
let row = rows
.into_iter()
.next()
.expect("rows.len() == 1 was just matched");
let decoded = T::from_pg_row(&row)?;
if let Some(target) = cache_target.as_ref() {
target.insert(&decoded).await;
}
Ok(decoded)
}
n => Err(DjogiError::multiple_objects(T::table_name(), n)),
}
}
}
pub fn fetch_all_prefetched<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<Vec<PrefetchedRow<T>>, DjogiError>> + Send + 'ctx
where
T::Pk: Clone + Send + Sync + 'static,
T: 'ctx,
{
async move {
if self.is_empty() {
return Ok(Vec::new());
}
auto_set_tenant::<T>(ctx).await?;
let prefetches = self.prefetch_paths.clone();
let acc = build_select(&self).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let pg_rows = ctx.query_all(&sql, ¶ms).await?;
let rows: Vec<T> = pg_rows
.iter()
.map(|r| T::from_pg_row(r))
.collect::<Result<Vec<T>, _>>()?;
apply_prefetches(ctx.inner_mut(), &prefetches, rows).await
}
}
pub fn fetch_all_joined<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<Vec<JoinedRow<T>>, DjogiError>> + Send + 'ctx
where
T: FromJoinedPgRow + 'ctx,
T::Pk: Clone + Send + Sync + 'static,
{
async move {
if self.is_empty() {
return Ok(Vec::new());
}
auto_set_tenant::<T>(ctx).await?;
let select_related_paths = self.select_related_paths.clone();
let prefetches = self.prefetch_paths.clone();
let acc = build_select_joined(&self).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows: Vec<tokio_postgres::Row> = ctx.query_all(&sql, ¶ms).await?;
let joined = apply_select_related::<T>(rows, &select_related_paths)?;
stitch_prefetches_into_joined(joined, &prefetches, ctx.inner_mut()).await
}
}
pub fn first<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<Option<T>, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
if self.is_empty() {
return Ok(None);
}
auto_set_tenant::<T>(ctx).await?;
let mut qs = self;
qs.limit = Some(1);
let cache_target = qs.cache_target.clone();
let acc = build_select(&qs).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let opt = ctx.query_opt(&sql, ¶ms).await?;
let decoded = opt.as_ref().map(|r| T::from_pg_row(r)).transpose()?;
if let (Some(target), Some(row)) = (cache_target.as_ref(), decoded.as_ref()) {
target.insert(row).await;
}
Ok(decoded)
}
}
pub fn get_or_create<'ctx, F>(
self,
ctx: &'ctx mut DjogiContext,
factory: F,
) -> impl Future<Output = Result<(T, bool), DjogiError>> + Send + 'ctx
where
F: FnOnce() -> T + Send + 'ctx,
T: 'ctx,
{
async move {
if let Some(row) = self.first(ctx).await? {
return Ok((row, false));
}
let created = T::create(ctx, factory()).await?;
Ok((created, true))
}
}
pub fn update_or_create<'ctx, F, U>(
self,
ctx: &'ctx mut DjogiContext,
factory: F,
updater: U,
) -> impl Future<Output = Result<(T, bool), DjogiError>> + Send + 'ctx
where
F: FnOnce() -> T + Send + 'ctx,
U: FnOnce(&mut T) + Send + 'ctx,
T: 'ctx,
{
async move {
if let Some(mut row) = self.first(ctx).await? {
updater(&mut row);
row.save(ctx).await?;
return Ok((row, false));
}
let created = T::create(ctx, factory()).await?;
Ok((created, true))
}
}
pub fn in_bulk<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
ids: Vec<T::Pk>,
) -> impl Future<Output = Result<HashMap<T::Pk, T>, DjogiError>> + Send + 'ctx
where
T::Pk: Eq + Hash,
T: 'ctx,
{
async move {
if self.is_empty() || ids.is_empty() {
return Ok(HashMap::new());
}
auto_set_tenant::<T>(ctx).await?;
let mut acc = SqlAccumulator::new("SELECT ");
acc.push_sql(<T as FromPgRow>::COLUMN_LIST);
acc.push_sql(" FROM ");
acc.push_sql(T::table_name());
acc.push_sql(" WHERE id IN (");
acc.push_list_binds(ids.iter().cloned());
acc.push_sql(")");
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let pg_rows = ctx.query_all(&sql, ¶ms).await?;
let mut out: HashMap<T::Pk, T> = HashMap::with_capacity(pg_rows.len());
for row in &pg_rows {
let item = T::from_pg_row(row)?;
out.insert(item.pk_value().clone(), item);
}
Ok(out)
}
}
}
impl<T: Model> QuerySet<T> {
pub fn count<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<i64, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
if self.is_empty() {
return Ok(0);
}
auto_set_tenant::<T>(ctx).await?;
let acc = build_count(&self).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let row = ctx.query_one(&sql, ¶ms).await?;
let n: i64 = try_get_scalar(&row, 0)?;
Ok(n)
}
}
pub fn exists<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<bool, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
if self.is_empty() {
return Ok(false);
}
auto_set_tenant::<T>(ctx).await?;
let acc = build_exists(&self).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let row = ctx.query_one(&sql, ¶ms).await?;
let b: bool = try_get_scalar(&row, 0)?;
Ok(b)
}
}
}
impl<T: Model> QuerySet<T>
where
T: FromPgRow + Unpin + Send,
{
pub async fn stream<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> Result<ModelCursorStream<'ctx, T>, DjogiError>
where
T: 'ctx,
{
self.stream_with_fetch_size(ctx, DEFAULT_FETCH_SIZE).await
}
pub async fn stream_with_fetch_size<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
fetch_size: u32,
) -> Result<ModelCursorStream<'ctx, T>, DjogiError>
where
T: 'ctx,
{
if fetch_size == 0 {
return Err(DjogiError::Validation(
"stream fetch_size must be at least 1".to_owned(),
));
}
auto_set_tenant::<T>(ctx).await?;
let acc = build_select(&self).map_err(DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
build_model_stream(ctx, &sql, ¶ms, fetch_size).await
}
}