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