use cratestack_core::{CratestackContext, CratestackError};
use crate::{ConflictTarget, ModelDescriptor, SqlColumnValue, SqlValue, SqlxRuntime, sqlx};
use super::upsert_do_nothing_sql::upsert_returning_record_do_nothing;
use super::upsert_do_update_sql::upsert_returning_record;
use super::upsert_sql::{row_passes_update_policy, select_for_update_by_conflict_target};
pub(super) struct UpsertResolution<M> {
pub(super) record: M,
pub(super) inserted: bool,
pub(super) before: Option<M>,
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn resolve_upsert<'tx, M, PK>(
tx: &mut sqlx::Transaction<'tx, sqlx::Postgres>,
runtime: &SqlxRuntime,
descriptor: &'static ModelDescriptor<M, PK>,
insert_values: &[SqlColumnValue],
conflict_target: ConflictTarget,
conflict_columns: &[(&'static str, SqlValue)],
ctx: &CratestackContext,
before_record: Option<M>,
) -> Result<UpsertResolution<M>, CratestackError>
where
for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow>,
{
if let Some(before) = before_record {
gate_update_policy(runtime, descriptor, conflict_columns, conflict_target, ctx).await?;
let record =
upsert_returning_record(&mut **tx, descriptor, insert_values, conflict_target).await?;
return Ok(UpsertResolution {
record,
inserted: false,
before: Some(before),
});
}
if let Some(record) =
upsert_returning_record_do_nothing(&mut **tx, descriptor, insert_values, conflict_target)
.await?
{
return Ok(UpsertResolution {
record,
inserted: true,
before: None,
});
}
let before = select_for_update_by_conflict_target(
&mut **tx,
descriptor,
conflict_columns,
conflict_target.predicate(),
)
.await?;
if before.is_some() {
gate_update_policy(runtime, descriptor, conflict_columns, conflict_target, ctx).await?;
}
let record =
upsert_returning_record(&mut **tx, descriptor, insert_values, conflict_target).await?;
let inserted = before.is_none();
Ok(UpsertResolution {
record,
inserted,
before,
})
}
async fn gate_update_policy<M, PK>(
runtime: &SqlxRuntime,
descriptor: &'static ModelDescriptor<M, PK>,
conflict_columns: &[(&'static str, SqlValue)],
conflict_target: ConflictTarget,
ctx: &CratestackContext,
) -> Result<(), CratestackError> {
if row_passes_update_policy(
runtime.pool(),
descriptor,
conflict_columns,
conflict_target.predicate(),
ctx,
)
.await?
{
return Ok(());
}
Err(CratestackError::Forbidden(
"update policy denied this upsert".to_owned(),
))
}