Skip to main content

cratestack_sqlx/query/write/
create_exec.rs

1//! Generic-over-Executor create helper used by both the pool and
2//! transaction paths in [`super::create`]. Validates, applies
3//! auth-defaults, seeds `@version`, evaluates create policies, then
4//! runs `INSERT ... RETURNING`.
5
6use cratestack_core::{CratestackContext, CratestackError};
7
8use crate::query::support::{
9    PolicyDb, apply_create_defaults, classify_unique_violation, evaluate_create_policies,
10    find_column_value, push_bind_value,
11};
12use crate::{CreateModelInput, ModelDescriptor, sqlx};
13
14/// Passing a pool as `policy_pool` while `executor` is a transaction takes a
15/// second connection; prefer the builder's `run_in_tx`.
16pub async fn create_record_with_executor<'e, E, M, PK, I>(
17    executor: E,
18    policy_pool: &sqlx::PgPool,
19    descriptor: &'static ModelDescriptor<M, PK>,
20    input: I,
21    ctx: &CratestackContext,
22) -> Result<M, CratestackError>
23where
24    E: sqlx::Executor<'e, Database = sqlx::Postgres>,
25    I: CreateModelInput<M>,
26    for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow> + serde::Serialize,
27{
28    let values =
29        authorized_create_values(PolicyDb::Pool(policy_pool), descriptor, input, ctx).await?;
30    insert_returning_record(executor, descriptor, &values).await
31}
32
33/// [`create_record_with_executor`] for a write that runs on `conn`, with
34/// the create-policy evaluation on `conn` too
35/// (docs/design/procedure-isolation.md ยง4.1).
36pub(crate) async fn create_record_in_conn<M, PK, I>(
37    conn: &mut sqlx::PgConnection,
38    descriptor: &'static ModelDescriptor<M, PK>,
39    input: I,
40    ctx: &CratestackContext,
41) -> Result<M, CratestackError>
42where
43    I: CreateModelInput<M>,
44    for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow> + serde::Serialize,
45{
46    let values =
47        authorized_create_values(PolicyDb::Conn(&mut *conn), descriptor, input, ctx).await?;
48    insert_returning_record(&mut *conn, descriptor, &values).await
49}
50
51async fn authorized_create_values<M, PK, I>(
52    policy: PolicyDb<'_>,
53    descriptor: &'static ModelDescriptor<M, PK>,
54    input: I,
55    ctx: &CratestackContext,
56) -> Result<Vec<crate::SqlColumnValue>, CratestackError>
57where
58    I: CreateModelInput<M>,
59{
60    input.validate()?;
61    let mut values = apply_create_defaults(input.sql_values(), descriptor.create_defaults, ctx)?;
62    // Seed the optimistic-lock column server-side. `@version` is
63    // excluded from the generated Create input so clients can't pick
64    // the initial value, and the column has no SQL `DEFAULT`. Done
65    // after `apply_create_defaults` so `@default`-driven overrides
66    // still win if a schema ever lands one.
67    if let Some(version_col) = descriptor.version_column
68        && find_column_value(&values, version_col).is_none()
69    {
70        values.push(crate::SqlColumnValue {
71            column: version_col,
72            value: crate::SqlValue::Int(0),
73        });
74    }
75    if values.is_empty() {
76        return Err(CratestackError::Validation(
77            "create input must contain at least one column".to_owned(),
78        ));
79    }
80    if !evaluate_create_policies(
81        policy,
82        descriptor.create_allow_policies,
83        descriptor.create_deny_policies,
84        &values,
85        ctx,
86    )
87    .await?
88    {
89        return Err(CratestackError::Forbidden(
90            "create policy denied this operation".to_owned(),
91        ));
92    }
93
94    Ok(values)
95}
96
97async fn insert_returning_record<'e, E, M, PK>(
98    executor: E,
99    descriptor: &'static ModelDescriptor<M, PK>,
100    values: &[crate::SqlColumnValue],
101) -> Result<M, CratestackError>
102where
103    E: sqlx::Executor<'e, Database = sqlx::Postgres>,
104    for<'r> M: Send + Unpin + sqlx::FromRow<'r, sqlx::postgres::PgRow>,
105{
106    let mut query = sqlx::QueryBuilder::<sqlx::Postgres>::new("INSERT INTO ");
107    query.push(descriptor.table_name).push(" (");
108    for (index, value) in values.iter().enumerate() {
109        if index > 0 {
110            query.push(", ");
111        }
112        query.push(value.column);
113    }
114    query.push(") VALUES (");
115    for (index, value) in values.iter().enumerate() {
116        if index > 0 {
117            query.push(", ");
118        }
119        push_bind_value(&mut query, &value.value);
120    }
121    query
122        .push(") RETURNING ")
123        .push(descriptor.select_projection());
124
125    query
126        .build_query_as::<M>()
127        .fetch_one(executor)
128        .await
129        .map_err(classify_unique_violation)
130}