#![allow(clippy::manual_async_fn)]
use crate::DjogiError;
use crate::context::DjogiContext;
use crate::model::Model;
use crate::pg::accumulator::as_params;
use crate::pg::decode::{FromJoinedPgRow, FromPgRow, decode_at};
use crate::query::condition::FilterValue;
use crate::query::field::{FieldRef, IntoFilterValue};
use crate::query::queryset::QuerySet;
use crate::query::returning::ReturningPair;
use crate::query::sql::{
build_delete, build_delete_returning, build_update, build_update_returning_ids,
build_update_returning_pairs,
};
use crate::query::terminal::auto_set_tenant;
use std::future::Future;
use std::marker::PhantomData;
#[derive(Debug, Clone)]
pub struct UpdateAssignment {
pub(crate) column: &'static str,
pub(crate) value: AssignmentValue,
}
#[derive(Debug, Clone)]
pub(crate) enum AssignmentValue {
Literal(FilterValue),
Expr(crate::expr::node::ExprNode),
}
impl UpdateAssignment {
pub(crate) fn new(column: &'static str, value: FilterValue) -> Self {
Self {
column,
value: AssignmentValue::Literal(value),
}
}
pub(crate) fn new_expr(column: &'static str, node: crate::expr::node::ExprNode) -> Self {
Self {
column,
value: AssignmentValue::Expr(node),
}
}
#[doc(hidden)]
pub fn column(&self) -> &'static str {
self.column
}
#[doc(hidden)]
pub(crate) fn value(&self) -> &AssignmentValue {
&self.value
}
}
impl<M: Model, V: IntoFilterValue> FieldRef<M, V> {
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn set(self, value: V) -> UpdateAssignment {
UpdateAssignment::new(self.column(), value.into_filter_value())
}
}
impl<M: Model, V: IntoFilterValue> FieldRef<M, V> {
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn set_expr(self, expr: crate::expr::Expr<V>) -> UpdateAssignment {
UpdateAssignment::new_expr(self.column(), expr.node)
}
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn set_field(self, other: FieldRef<M, V>) -> UpdateAssignment {
UpdateAssignment::new_expr(self.column(), other.as_expr().node)
}
}
impl<M: Model, V> FieldRef<M, V>
where
V: IntoFilterValue + crate::expr::arithmetic::Numeric + Into<crate::expr::Expr<V>>,
{
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn increment(self, amount: V) -> UpdateAssignment {
let expr = self.as_expr() + crate::expr::Expr::literal(amount);
UpdateAssignment::new_expr(self.column(), expr.node)
}
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn decrement(self, amount: V) -> UpdateAssignment {
let expr = self.as_expr() - crate::expr::Expr::literal(amount);
UpdateAssignment::new_expr(self.column(), expr.node)
}
}
impl<M: Model> FieldRef<M, crate::Interval> {
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn increment(self, amount: crate::Interval) -> UpdateAssignment {
let expr = self.as_expr() + crate::expr::Expr::literal(amount);
UpdateAssignment::new_expr(self.column(), expr.node)
}
#[must_use = "assignments are lazy — drop one and the SET clause is silently omitted"]
pub fn decrement(self, amount: crate::Interval) -> UpdateAssignment {
let expr = self.as_expr() - crate::expr::Expr::literal(amount);
UpdateAssignment::new_expr(self.column(), expr.node)
}
}
pub trait IntoAssignments {
fn into_assignments(self) -> Vec<UpdateAssignment>;
}
impl IntoAssignments for UpdateAssignment {
fn into_assignments(self) -> Vec<UpdateAssignment> {
vec![self]
}
}
impl IntoAssignments for Vec<UpdateAssignment> {
fn into_assignments(self) -> Vec<UpdateAssignment> {
self
}
}
#[must_use = "UpdateStmt is inert — call .execute(ctx) to run the UPDATE"]
pub struct UpdateStmt<T: Model> {
pub(crate) qs: QuerySet<T>,
pub(crate) assignments: Vec<UpdateAssignment>,
pub(crate) _m: PhantomData<fn() -> T>,
}
impl<T: Model> Clone for UpdateStmt<T> {
fn clone(&self) -> Self {
UpdateStmt {
qs: self.qs.clone(),
assignments: self.assignments.clone(),
_m: PhantomData,
}
}
}
impl<T: Model> std::fmt::Debug for UpdateStmt<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UpdateStmt")
.field("table", &T::table_name())
.field("qs", &self.qs)
.field("assignments", &self.assignments)
.finish()
}
}
impl<T: Model> UpdateStmt<T> {
pub fn execute<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<u64, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
self.qs.validate_mutation_read_tail("update")?;
if self.qs.is_empty() || self.assignments.is_empty() {
return Ok(0);
}
auto_set_tenant::<T>(ctx).await?;
if T::__djogi_should_collect_bulk_update_ids(ctx) {
let acc = build_update_returning_ids(&self.qs, &self.assignments)
.map_err(crate::DjogiError::from)?;
let pk_column = T::descriptor().pk_column().expect(
"Model implementations with CRUD support must expose a primary-key column",
);
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows = ctx.query_all(&sql, ¶ms).await?;
let mut ids = Vec::with_capacity(rows.len());
for row in &rows {
ids.push(decode_at::<T::Pk>(row, 0, pk_column)?);
}
let rows_affected = ids.len() as u64;
<T as Model>::__djogi_enqueue_bulk_on_save_cache_invalidation(ctx, ids)?;
Ok(rows_affected)
} else {
let acc =
build_update(&self.qs, &self.assignments).map_err(crate::DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows_affected = ctx.execute(&sql, ¶ms).await?;
Ok(rows_affected)
}
}
}
pub fn execute_returning_pairs<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<Vec<ReturningPair<T>>, DjogiError>> + Send + 'ctx
where
T: FromPgRow + FromJoinedPgRow + 'ctx,
{
async move {
self.qs
.validate_mutation_read_tail("execute_returning_pairs")?;
if self.qs.is_empty() || self.assignments.is_empty() {
return Ok(Vec::new());
}
auto_set_tenant::<T>(ctx).await?;
let acc = build_update_returning_pairs(&self.qs, &self.assignments)
.map_err(crate::DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows = ctx.query_all(&sql, ¶ms).await?;
let mut pairs = Vec::with_capacity(rows.len());
for row in &rows {
let old = T::from_joined_pg_row(row, "__djogi_old__")?;
let new = T::from_joined_pg_row(row, "__djogi_new__")?;
pairs.push(ReturningPair { old, new });
}
let outbox_rows: Vec<&T> = pairs.iter().map(|pair| &pair.new).collect();
<T as Model>::__djogi_emit_save_outbox_batch(ctx, &outbox_rows).await?;
let cache_ids: Vec<T::Pk> = pairs
.iter()
.map(|pair| pair.new.pk_value().clone())
.collect();
<T as Model>::__djogi_enqueue_bulk_on_save_cache_invalidation(ctx, cache_ids)?;
Ok(pairs)
}
}
}
impl<T: Model> QuerySet<T> {
#[must_use = "UpdateStmt is inert — call .execute(ctx) to run the UPDATE"]
pub fn update<F, A>(self, f: F) -> UpdateStmt<T>
where
F: FnOnce(T::Fields) -> A,
A: IntoAssignments,
{
let assignments = f(T::Fields::default()).into_assignments();
UpdateStmt {
qs: self,
assignments,
_m: PhantomData,
}
}
pub fn delete<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<u64, DjogiError>> + Send + 'ctx
where
T: 'ctx,
{
async move {
self.validate_mutation_read_tail("delete")?;
if self.is_empty() {
return Ok(0);
}
auto_set_tenant::<T>(ctx).await?;
let acc = build_delete(&self).map_err(crate::DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows_affected = ctx.execute(&sql, ¶ms).await?;
Ok(rows_affected)
}
}
pub fn delete_returning<'ctx>(
self,
ctx: &'ctx mut DjogiContext,
) -> impl Future<Output = Result<Vec<T>, DjogiError>> + Send + 'ctx
where
T: FromPgRow + FromJoinedPgRow + 'ctx,
{
async move {
self.validate_mutation_read_tail("delete_returning")?;
if self.is_empty() {
return Ok(Vec::new());
}
auto_set_tenant::<T>(ctx).await?;
let acc = build_delete_returning(&self).map_err(crate::DjogiError::from)?;
let (sql, binds) = acc.into_parts();
let params = as_params(&binds);
let rows = ctx.query_all(&sql, ¶ms).await?;
let mut deleted = Vec::with_capacity(rows.len());
for row in &rows {
deleted.push(T::from_joined_pg_row(row, "__djogi_old__")?);
}
Ok(deleted)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::descriptor::ModelDescriptor;
use crate::query::field::FieldRef;
struct Fake;
impl crate::model::__sealed::Sealed for Fake {}
#[allow(clippy::manual_async_fn)]
impl Model for Fake {
type Pk = i64;
type Fields = ();
fn table_name() -> &'static str {
"fakes"
}
fn pk_value(&self) -> &i64 {
unreachable!()
}
fn descriptor() -> &'static ModelDescriptor {
unreachable!()
}
fn get(
_ctx: &mut crate::context::DjogiContext,
_id: i64,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn create(
_ctx: &mut crate::context::DjogiContext,
_v: Self,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send {
async { unreachable!() }
}
fn save<'ctx>(
&'ctx mut self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
fn delete(
self,
_ctx: &mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<(), crate::DjogiError>> + Send {
async { unreachable!() }
}
fn refresh_from_db<'ctx>(
&'ctx self,
_ctx: &'ctx mut crate::context::DjogiContext,
) -> impl std::future::Future<Output = Result<Self, crate::DjogiError>> + Send + 'ctx
{
async { unreachable!() }
}
}
#[test]
fn field_ref_set_builds_assignment_with_projected_value() {
let f: FieldRef<Fake, i32> = FieldRef::new("view_count");
let a = f.set(42i32);
assert_eq!(a.column(), "view_count");
assert!(matches!(
a.value(),
crate::query::update::AssignmentValue::Literal(FilterValue::I32(42))
));
}
#[test]
fn field_ref_set_expr_builds_expression_assignment() {
use crate::expr::Expr;
let f: FieldRef<Fake, i64> = FieldRef::new("balance");
let a = f.set_expr(f.as_expr() + Expr::literal(5i64));
assert_eq!(a.column(), "balance");
assert!(matches!(
a.value(),
crate::query::update::AssignmentValue::Expr(_)
));
}
#[test]
fn into_assignments_single_wraps_in_vec() {
let f: FieldRef<Fake, bool> = FieldRef::new("published");
let a = f.set(true);
let v = a.into_assignments();
assert_eq!(v.len(), 1);
assert_eq!(v[0].column(), "published");
}
#[test]
fn into_assignments_vec_passes_through() {
let a: FieldRef<Fake, i32> = FieldRef::new("view_count");
let b: FieldRef<Fake, bool> = FieldRef::new("published");
let vs = vec![a.set(0i32), b.set(false)];
let out = vs.into_assignments();
assert_eq!(out.len(), 2);
assert_eq!(out[0].column(), "view_count");
assert_eq!(out[1].column(), "published");
}
#[test]
fn update_stmt_clones_preserve_assignments() {
let f: FieldRef<Fake, i32> = FieldRef::new("view_count");
let qs: QuerySet<Fake> = QuerySet::new();
let stmt = qs.update(|_| f.set(42i32));
let cloned = stmt.clone();
assert_eq!(cloned.assignments.len(), 1);
assert_eq!(cloned.assignments[0].column(), "view_count");
}
#[test]
fn set_field_builds_expr_assignment() {
let target: FieldRef<Fake, i64> = FieldRef::new("balance");
let source: FieldRef<Fake, i64> = FieldRef::new("overdraft_limit");
let a = target.set_field(source);
assert_eq!(a.column(), "balance");
assert!(
matches!(a.value(), AssignmentValue::Expr(_)),
"set_field must produce an Expr assignment, not a Literal"
);
}
#[test]
fn increment_builds_add_expr_assignment() {
let f: FieldRef<Fake, i64> = FieldRef::new("balance");
let a = f.increment(10i64);
assert_eq!(a.column(), "balance");
assert!(
matches!(a.value(), AssignmentValue::Expr(_)),
"increment must produce an Expr assignment"
);
}
}