cratestack_sqlx/query/write/
update_exec.rs1use 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
42pub(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}