use drizzle_core::error::DrizzleError;
use drizzle_core::traits::ToSQL;
use drizzle_postgres::builder::{DeleteInitial, InsertInitial, SelectInitial, UpdateInitial};
use drizzle_postgres::traits::PostgresTable;
use std::marker::PhantomData;
use tokio_postgres::{Row, Transaction as TokioPgTransaction};
use crate::builder::postgres::tokio_postgres::{Rows, prepared::ClientStatementCache};
use crate::transaction::savepoint::{AsyncSavepointState, async_savepoint};
#[cfg(feature = "query")]
use crate::builder::postgres::common;
fn tx_consumed_error() -> DrizzleError {
DrizzleError::TransactionError("Transaction already consumed".into())
}
use drizzle_postgres::builder::{
self, QueryBuilder, delete::DeleteBuilder, insert::InsertBuilder, select::SelectBuilder,
update::UpdateBuilder,
};
use drizzle_postgres::common::PostgresTransactionType;
use drizzle_postgres::transaction::{IsolationLevel, TransactionConfig};
use drizzle_postgres::values::PostgresValue;
use crate::builder::postgres::tokio_postgres::tokio_postgres_materialize_params as materialize_params;
pub type TransactionBuilder<'tx, 'conn, Schema, Builder, State> =
crate::transaction::postgres::typestate::TransactionBuilder<
'tx,
&'tx Transaction<'conn, Schema>,
Schema,
Builder,
State,
>;
use crate::builder::postgres::tokio_postgres::prepared;
use drizzle_core::prepared::prepare_render;
crate::drizzle_tx_prepare_impl!('conn);
pub struct Transaction<'conn, Schema = ()> {
tx: Option<TokioPgTransaction<'conn>>,
config: TransactionConfig,
savepoints: AsyncSavepointState,
schema: Schema,
statement_cache: ClientStatementCache,
}
impl<Schema> std::fmt::Debug for Transaction<'_, Schema> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Transaction")
.field("config", &self.config)
.field("is_active", &self.tx.is_some())
.finish()
}
}
impl<'conn, Schema> Transaction<'conn, Schema> {
pub(crate) fn new(
tx: TokioPgTransaction<'conn>,
config: TransactionConfig,
schema: Schema,
statement_cache: ClientStatementCache,
) -> Self {
Self {
tx: Some(tx),
config,
savepoints: AsyncSavepointState::new(),
schema,
statement_cache,
}
}
#[inline]
pub const fn schema(&self) -> &Schema {
&self.schema
}
#[deprecated(since = "0.2.0", note = "use config()")]
#[inline]
pub const fn tx_type(&self) -> PostgresTransactionType {
match self.config.isolation() {
None | Some(IsolationLevel::ReadCommitted) => PostgresTransactionType::ReadCommitted,
Some(IsolationLevel::ReadUncommitted) => PostgresTransactionType::ReadUncommitted,
Some(IsolationLevel::RepeatableRead) => PostgresTransactionType::RepeatableRead,
Some(IsolationLevel::Serializable) => PostgresTransactionType::Serializable,
}
}
#[inline]
pub const fn config(&self) -> TransactionConfig {
self.config
}
fn statement_error(&self, error: tokio_postgres::Error) -> DrizzleError {
if error.as_db_error().is_some() {
self.savepoints.aborted().mark();
}
if crate::builder::postgres::tokio_postgres::prepared::is_stale_statement(&error) {
self.statement_cache.clear();
}
DrizzleError::from(error)
}
async fn execute_raw(&self, sql: &str) -> drizzle_core::error::Result<()> {
self.savepoints.ensure_usable()?;
let tx = self.tx.as_ref().ok_or_else(tx_consumed_error)?;
tx.execute(sql, &[]).await.map_err(DrizzleError::from)?;
Ok(())
}
pub async fn savepoint<F, R>(&self, f: F) -> drizzle_core::error::Result<R>
where
F: AsyncFnOnce(&Self) -> drizzle_core::error::Result<R>,
{
async_savepoint(
&self.savepoints,
|sql| async move { self.execute_raw(&sql).await },
f(self),
)
.await
}
postgres_transaction_constructors!('conn);
pub async fn execute<'q, T>(&self, query: T) -> drizzle_core::error::Result<u64>
where
T: ToSQL<'q, PostgresValue<'q>>,
{
self.savepoints.ensure_usable()?;
let query_sql = query.to_sql();
let (sql, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx.execute");
let (sql, params) = query_sql.build();
drizzle_core::drizzle_trace_query!(&sql, params.len());
(sql, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.statement_cache
.transaction_statement(tx, &sql, ¶m_types)
.await
.map_err(|error| self.statement_error(error))?;
Ok(tx
.execute(&statement, ¶m_refs[..])
.await
.map_err(|error| self.statement_error(error))?)
}
pub async fn all<'q, T, R, C>(&self, query: T) -> drizzle_core::error::Result<C>
where
R: for<'r> TryFrom<&'r Row>,
for<'r> <R as TryFrom<&'r Row>>::Error: Into<drizzle_core::error::DrizzleError>,
T: ToSQL<'q, PostgresValue<'q>>,
C: std::iter::FromIterator<R>,
{
self.savepoints.ensure_usable()?;
let sql = query.to_sql();
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx.all");
let (sql_str, params) = sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.statement_error(error))?;
let rows = tx
.query(&statement, ¶m_refs[..])
.await
.map_err(|error| self.statement_error(error))?;
let mut decoded = Vec::with_capacity(rows.len());
for row in rows {
decoded.push(R::try_from(&row).map_err(Into::into)?);
}
Ok(decoded.into_iter().collect())
}
pub async fn rows<'q, T, R>(&self, query: T) -> drizzle_core::error::Result<Rows<R>>
where
R: for<'r> TryFrom<&'r Row>,
for<'r> <R as TryFrom<&'r Row>>::Error: Into<drizzle_core::error::DrizzleError>,
T: ToSQL<'q, PostgresValue<'q>>,
{
self.savepoints.ensure_usable()?;
let sql = query.to_sql();
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx.rows");
let (sql_str, params) = sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.statement_error(error))?;
let rows = tx
.query(&statement, ¶m_refs[..])
.await
.map_err(|error| self.statement_error(error))?;
Ok(Rows::new(rows))
}
pub async fn get<'q, T, R>(&self, query: T) -> drizzle_core::error::Result<R>
where
R: for<'r> TryFrom<&'r Row>,
for<'r> <R as TryFrom<&'r Row>>::Error: Into<drizzle_core::error::DrizzleError>,
T: ToSQL<'q, PostgresValue<'q>>,
{
self.savepoints.ensure_usable()?;
let sql = query.to_sql();
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx.get");
let (sql_str, params) = sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.statement_error(error))?;
let row = tx
.query_one(&statement, ¶m_refs[..])
.await
.map_err(|error| self.statement_error(error))?;
R::try_from(&row).map_err(Into::into)
}
#[cfg(feature = "query")]
pub fn query<'a, T>(&self, _table: T) -> common::DrizzleQueryBuilder<'_, 'a, &Self, Schema, T>
where
T: drizzle_core::query::QueryTable,
{
common::DrizzleQueryBuilder {
runner: self,
builder: drizzle_core::query::QueryBuilder::new(),
_schema: PhantomData,
}
}
pub(crate) async fn commit(mut self) -> drizzle_core::error::Result<()> {
if let Err(error) = self.savepoints.ensure_usable() {
let tx = self.tx.take().ok_or_else(tx_consumed_error)?;
tx.rollback().await.map_err(DrizzleError::from)?;
return Err(error);
}
if self.savepoints.aborted().is_aborted() {
let tx = self.tx.take().ok_or_else(tx_consumed_error)?;
tx.rollback().await.map_err(DrizzleError::from)?;
return Err(crate::transaction::savepoint::aborted_transaction_error());
}
let tx = self.tx.take().ok_or_else(tx_consumed_error)?;
tx.commit().await.map_err(DrizzleError::from)
}
pub(crate) async fn rollback(mut self) -> drizzle_core::error::Result<()> {
let tx = self.tx.take().ok_or_else(tx_consumed_error)?;
tx.rollback().await.map_err(DrizzleError::from)
}
}
#[cfg(feature = "query")]
impl<Schema> common::RelationalPreparedDriver for &Transaction<'_, Schema> {
type PreparedDriver = tokio_postgres::Client;
}
#[cfg(feature = "query")]
use drizzle_core::query::{DeserializeStore, FromJsonObject as _};
#[cfg(feature = "query")]
impl<'db, 'a, 'conn, Schema, T, Rels, Cl>
common::DrizzleQueryBuilder<
'db,
'a,
&'db Transaction<'conn, Schema>,
Schema,
T,
Rels,
drizzle_core::query::AllColumns,
Cl,
>
{
pub async fn find_many(
self,
) -> drizzle_core::error::Result<
Vec<
<Rels as drizzle_core::query::BuildRow<
<T as drizzle_core::query::QueryTable>::Select,
>>::Row,
>,
>
where
T: drizzle_core::query::QueryTable,
<T as drizzle_core::query::QueryTable>::Select: for<'r> TryFrom<&'r Row>,
for<'r> <<T as drizzle_core::query::QueryTable>::Select as TryFrom<&'r Row>>::Error:
Into<DrizzleError>,
Rels: drizzle_core::query::BuildRow<<T as drizzle_core::query::QueryTable>::Select>
+ drizzle_core::query::RenderRelations<'a, PostgresValue<'a>>,
<Rels as drizzle_core::query::BuildStore>::Store: DeserializeStore,
{
self.runner.savepoints.ensure_usable()?;
let num_base_cols = T::COLUMN_NAMES.len();
let builder = self.builder;
let mut rendered = Vec::new();
builder.relations.render_into(&mut rendered);
let query_sql = drizzle_core::query::build_query_sql(
T::TABLE,
T::COLUMN_NAMES,
T::BLOB_COLUMNS,
T::JSON_PROJECTIONS,
rendered,
builder.where_sql,
builder.order_by_sql,
builder.limit,
builder.offset,
false,
);
let (sql, bind_params) = query_sql.build();
drizzle_core::drizzle_trace_query!(&sql, bind_params.len());
let (param_types, param_refs) = materialize_params(&bind_params);
let tx = self.runner.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.runner
.statement_cache
.transaction_statement(tx, &sql, ¶m_types)
.await
.map_err(|error| self.runner.statement_error(error))?;
let rows = tx
.query(&statement, ¶m_refs[..])
.await
.map_err(|error| self.runner.statement_error(error))?;
let mut results = Vec::with_capacity(rows.len());
for row in &rows {
let base = <T as drizzle_core::query::QueryTable>::Select::try_from(row)
.map_err(Into::into)?;
let mut rel_col = num_base_cols;
let mut next_rel = || {
let json: Option<String> = row.get(rel_col);
rel_col += 1;
Ok(json)
};
let store =
<Rels as drizzle_core::query::BuildStore>::Store::from_json_columns(&mut next_rel)?;
results.push(<Rels as drizzle_core::query::BuildRow<_>>::assemble(
base, store,
));
}
Ok(results)
}
}
#[cfg(feature = "query")]
impl<'db, 'a, 'conn, Schema, T, Rels, W, Ord>
common::DrizzleQueryBuilder<
'db,
'a,
&'db Transaction<'conn, Schema>,
Schema,
T,
Rels,
drizzle_core::query::AllColumns,
drizzle_core::query::Clauses<W, Ord, drizzle_core::query::NoLimit>,
>
{
pub async fn find_first(
self,
) -> drizzle_core::error::Result<
Option<
<Rels as drizzle_core::query::BuildRow<
<T as drizzle_core::query::QueryTable>::Select,
>>::Row,
>,
>
where
T: drizzle_core::query::QueryTable,
<T as drizzle_core::query::QueryTable>::Select: for<'r> TryFrom<&'r Row>,
for<'r> <<T as drizzle_core::query::QueryTable>::Select as TryFrom<&'r Row>>::Error:
Into<DrizzleError>,
Rels: drizzle_core::query::BuildRow<<T as drizzle_core::query::QueryTable>::Select>
+ drizzle_core::query::RenderRelations<'a, PostgresValue<'a>>,
<Rels as drizzle_core::query::BuildStore>::Store: DeserializeStore,
{
Ok(self.limit(1).find_many().await?.into_iter().next())
}
}
#[cfg(feature = "query")]
impl<'db, 'a, 'conn, Schema, T, Rels, Cl>
common::DrizzleQueryBuilder<
'db,
'a,
&'db Transaction<'conn, Schema>,
Schema,
T,
Rels,
drizzle_core::query::PartialColumns,
Cl,
>
{
pub async fn find_many(
self,
) -> drizzle_core::error::Result<
Vec<
<Rels as drizzle_core::query::BuildRow<
<T as drizzle_core::query::QueryTable>::PartialSelect,
>>::Row,
>,
>
where
T: drizzle_core::query::QueryTable,
<T as drizzle_core::query::QueryTable>::PartialSelect: drizzle_core::query::FromJsonObject,
Rels: drizzle_core::query::BuildRow<<T as drizzle_core::query::QueryTable>::PartialSelect>
+ drizzle_core::query::RenderRelations<'a, PostgresValue<'a>>,
<Rels as drizzle_core::query::BuildStore>::Store: DeserializeStore,
{
self.runner.savepoints.ensure_usable()?;
let builder = self.builder;
let column_names = &builder.cols.columns;
let mut rendered = Vec::new();
builder.relations.render_into(&mut rendered);
let col_refs: Vec<&str> = column_names.clone();
let query_sql = drizzle_core::query::build_query_sql(
T::TABLE,
&col_refs,
T::BLOB_COLUMNS,
T::JSON_PROJECTIONS,
rendered,
builder.where_sql,
builder.order_by_sql,
builder.limit,
builder.offset,
true,
);
let (sql, bind_params) = query_sql.build();
drizzle_core::drizzle_trace_query!(&sql, bind_params.len());
let (param_types, param_refs) = materialize_params(&bind_params);
let tx = self.runner.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.runner
.statement_cache
.transaction_statement(tx, &sql, ¶m_types)
.await
.map_err(|error| self.runner.statement_error(error))?;
let rows = tx
.query(&statement, ¶m_refs[..])
.await
.map_err(|error| self.runner.statement_error(error))?;
let mut results = Vec::with_capacity(rows.len());
for row in &rows {
let base_json: String = row.get(0);
let base = <T as drizzle_core::query::QueryTable>::PartialSelect::from_json_str(
&base_json, "base",
)?;
let mut rel_col = 1usize;
let mut next_rel = || {
let json: Option<String> = row.get(rel_col);
rel_col += 1;
Ok(json)
};
let store =
<Rels as drizzle_core::query::BuildStore>::Store::from_json_columns(&mut next_rel)?;
results.push(<Rels as drizzle_core::query::BuildRow<_>>::assemble(
base, store,
));
}
Ok(results)
}
}
#[cfg(feature = "query")]
impl<'db, 'a, 'conn, Schema, T, Rels, W, Ord>
common::DrizzleQueryBuilder<
'db,
'a,
&'db Transaction<'conn, Schema>,
Schema,
T,
Rels,
drizzle_core::query::PartialColumns,
drizzle_core::query::Clauses<W, Ord, drizzle_core::query::NoLimit>,
>
{
pub async fn find_first(
self,
) -> drizzle_core::error::Result<
Option<
<Rels as drizzle_core::query::BuildRow<
<T as drizzle_core::query::QueryTable>::PartialSelect,
>>::Row,
>,
>
where
T: drizzle_core::query::QueryTable,
<T as drizzle_core::query::QueryTable>::PartialSelect: drizzle_core::query::FromJsonObject,
Rels: drizzle_core::query::BuildRow<<T as drizzle_core::query::QueryTable>::PartialSelect>
+ drizzle_core::query::RenderRelations<'a, PostgresValue<'a>>,
<Rels as drizzle_core::query::BuildStore>::Store: DeserializeStore,
{
Ok(self.limit(1).find_many().await?.into_iter().next())
}
}
impl<'tx, 'q, S, Schema, State, Table, Mk, Rw, Grouped>
TransactionBuilder<'tx, '_, S, QueryBuilder<'q, Schema, State, Table, Mk, Rw, Grouped>, State>
where
State: builder::ExecutableState,
{
pub async fn execute(self) -> drizzle_core::error::Result<u64> {
self.runner.savepoints.ensure_usable()?;
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx_builder.execute");
let (sql_str, params) = self.builder.sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.runner.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.runner
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.runner.statement_error(error))?;
Ok(tx
.execute(&statement, ¶m_refs[..])
.await
.map_err(|error| self.runner.statement_error(error))?)
}
pub async fn all<R, Proof, AggProof>(self) -> drizzle_core::error::Result<Vec<R>>
where
for<'r> Mk: drizzle_core::row::DecodeSelectedRef<&'r ::tokio_postgres::Row, R>
+ drizzle_core::row::MarkerScopeValidFor<Proof>
+ drizzle_core::row::StrictDecodeMarker
+ drizzle_core::row::MarkerColumnCountValid<::tokio_postgres::Row, Rw, R>,
Mk: drizzle_core::row::MarkerAggValidFor<Grouped, AggProof>,
{
self.runner.savepoints.ensure_usable()?;
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx_builder.all");
let (sql_str, params) = self.builder.sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.runner.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.runner
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.runner.statement_error(error))?;
let rows = tx
.query(&statement, ¶m_refs[..])
.await
.map_err(|error| self.runner.statement_error(error))?;
let mut decoded = Vec::with_capacity(rows.len());
for row in &rows {
decoded.push(<Mk as drizzle_core::row::DecodeSelectedRef<
&::tokio_postgres::Row,
R,
>>::decode(row)?);
}
Ok(decoded)
}
pub async fn rows(self) -> drizzle_core::error::Result<Rows<Rw>>
where
Rw: for<'r> TryFrom<&'r Row>,
for<'r> <Rw as TryFrom<&'r Row>>::Error: Into<drizzle_core::error::DrizzleError>,
{
self.runner.savepoints.ensure_usable()?;
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx_builder.rows");
let (sql_str, params) = self.builder.sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.runner.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.runner
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.runner.statement_error(error))?;
let rows = tx
.query(&statement, ¶m_refs[..])
.await
.map_err(|error| self.runner.statement_error(error))?;
Ok(Rows::new(rows))
}
pub async fn get<R, Proof, AggProof>(self) -> drizzle_core::error::Result<R>
where
for<'r> Mk: drizzle_core::row::DecodeSelectedRef<&'r ::tokio_postgres::Row, R>
+ drizzle_core::row::MarkerScopeValidFor<Proof>
+ drizzle_core::row::StrictDecodeMarker
+ drizzle_core::row::MarkerColumnCountValid<::tokio_postgres::Row, Rw, R>,
Mk: drizzle_core::row::MarkerAggValidFor<Grouped, AggProof>,
{
self.runner.savepoints.ensure_usable()?;
let (sql_str, params) = {
#[cfg(feature = "profiling")]
drizzle_core::drizzle_profile_scope!("postgres.tokio", "tx_builder.get");
let (sql_str, params) = self.builder.sql.build();
drizzle_core::drizzle_trace_query!(&sql_str, params.len());
(sql_str, params)
};
let (param_types, param_refs) = materialize_params(¶ms);
let tx = self.runner.tx.as_ref().ok_or_else(tx_consumed_error)?;
let statement = self
.runner
.statement_cache
.transaction_statement(tx, &sql_str, ¶m_types)
.await
.map_err(|error| self.runner.statement_error(error))?;
let row = tx
.query_one(&statement, ¶m_refs[..])
.await
.map_err(|error| self.runner.statement_error(error))?;
<Mk as drizzle_core::row::DecodeSelectedRef<&::tokio_postgres::Row, R>>::decode(&row)
}
}