Skip to main content

cratestack_sqlx/query/write/
update_exec.rs

1//! Generic-over-Executor update helpers used by single-row UPDATE
2//! paths. Builds `UPDATE ... SET ... WHERE pk = $X [AND version = $Y]
3//! AND policy(...) RETURNING ...`, with version-mismatch detection via
4//! a read-policy probe.
5//!
6//! Two entry points, differing only in where that probe runs:
7//! [`update_record_with_executor`] (public; the probe runs on the pool it
8//! is handed) and [`update_record_in_conn`] (the probe runs on the
9//! statement's own connection, for every caller; see
10//! docs/design/procedure-isolation.md ยง4.1).
11
12use cratestack_core::{CratestackContext, CratestackError};
13
14use crate::query::support::{
15    classify_unique_violation, no_row_error, push_action_policy_query, push_bind_value,
16};
17use crate::{ModelDescriptor, SqlColumnValue, UpdateModelInput, sqlx};
18
19/// Passing a pool as `policy_pool` while `executor` is a transaction takes a
20/// second connection; prefer the builder's `run_in_tx`.
21pub async fn update_record_with_executor<'e, E, M, PK, I>(
22    executor: E,
23    policy_pool: &sqlx::PgPool,
24    descriptor: &'static ModelDescriptor<M, PK>,
25    id: PK,
26    input: I,
27    ctx: &CratestackContext,
28    if_match: Option<i64>,
29) -> Result<M, CratestackError>
30where
31    E: sqlx::Executor<'e, Database = sqlx::Postgres>,
32    I: UpdateModelInput<M>,
33    for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow> + serde::Serialize,
34    PK: Send + Clone + sqlx::Type<sqlx::Postgres> + for<'q> sqlx::Encode<'q, sqlx::Postgres>,
35{
36    let values = update_values(input)?;
37    let probe_id = id.clone();
38    match update_returning_record(executor, descriptor, id, &values, ctx, if_match).await? {
39        Some(record) => Ok(record),
40        None => Err(no_row_error(policy_pool, descriptor, probe_id, ctx, if_match, "update").await),
41    }
42}
43
44/// [`update_record_with_executor`] for a write that runs on `conn`, with
45/// the version/policy probe on `conn` too.
46pub(crate) async fn update_record_in_conn<M, PK, I>(
47    conn: &mut sqlx::PgConnection,
48    descriptor: &'static ModelDescriptor<M, PK>,
49    id: PK,
50    input: I,
51    ctx: &CratestackContext,
52    if_match: Option<i64>,
53) -> Result<M, CratestackError>
54where
55    I: UpdateModelInput<M>,
56    for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow> + serde::Serialize,
57    PK: Send + Clone + sqlx::Type<sqlx::Postgres> + for<'q> sqlx::Encode<'q, sqlx::Postgres>,
58{
59    let values = update_values(input)?;
60    let probe_id = id.clone();
61    match update_returning_record(&mut *conn, descriptor, id, &values, ctx, if_match).await? {
62        Some(record) => Ok(record),
63        None => Err(no_row_error(conn, descriptor, probe_id, ctx, if_match, "update").await),
64    }
65}
66
67fn update_values<M, I: UpdateModelInput<M>>(
68    input: I,
69) -> Result<Vec<SqlColumnValue>, CratestackError> {
70    input.validate()?;
71    let values = input.sql_values();
72    if values.is_empty() {
73        return Err(CratestackError::Validation(
74            "update input must contain at least one changed column".to_owned(),
75        ));
76    }
77    Ok(values)
78}
79
80async fn update_returning_record<'e, E, M, PK>(
81    executor: E,
82    descriptor: &'static ModelDescriptor<M, PK>,
83    id: PK,
84    values: &[crate::SqlColumnValue],
85    ctx: &CratestackContext,
86    if_match: Option<i64>,
87) -> Result<Option<M>, CratestackError>
88where
89    E: sqlx::Executor<'e, Database = sqlx::Postgres>,
90    for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow>,
91    PK: Send + Clone + sqlx::Type<sqlx::Postgres> + for<'q> sqlx::Encode<'q, sqlx::Postgres>,
92{
93    let version_column = descriptor.version_column;
94    let mut query = sqlx::QueryBuilder::<sqlx::Postgres>::new("UPDATE ");
95    query.push(descriptor.table_name).push(" SET ");
96    for (index, value) in values.iter().enumerate() {
97        if index > 0 {
98            query.push(", ");
99        }
100        query.push(value.column).push(" = ");
101        push_bind_value(&mut query, &value.value);
102    }
103    if let Some(version_col) = version_column {
104        query
105            .push(", ")
106            .push(version_col)
107            .push(" = ")
108            .push(version_col)
109            .push(" + 1");
110    }
111    query
112        .push(" WHERE ")
113        .push(descriptor.primary_key)
114        .push(" = ");
115    query.push_bind(id);
116    if let (Some(version_col), Some(expected)) = (version_column, if_match) {
117        query.push(" AND ").push(version_col).push(" = ");
118        query.push_bind(expected);
119    }
120    query.push(" AND ");
121    push_action_policy_query(
122        &mut query,
123        descriptor.update_allow_policies,
124        descriptor.update_deny_policies,
125        ctx,
126    );
127    query
128        .push(" RETURNING ")
129        .push(descriptor.select_projection());
130
131    query
132        .build_query_as::<M>()
133        .fetch_optional(executor)
134        .await
135        .map_err(classify_unique_violation)
136}