use drizzle_core::error::DrizzleError;
use drizzle_core::traits::ToSQL;
#[cfg(feature = "sqlite")]
use drizzle_sqlite::builder::{DeleteInitial, InsertInitial, SelectInitial, UpdateInitial};
#[cfg(feature = "sqlite")]
use drizzle_sqlite::traits::SQLiteTable;
use libsql::Row;
use std::marker::PhantomData;
use std::sync::atomic::AtomicU32;
use crate::builder::sqlite::rows::LibsqlRows as Rows;
use crate::transaction::savepoint::async_savepoint;
#[cfg(feature = "sqlite")]
use drizzle_sqlite::{
builder::{
self, QueryBuilder, delete::DeleteBuilder, insert::InsertBuilder, select::SelectBuilder,
update::UpdateBuilder,
},
connection::SQLiteTransactionType,
values::SQLiteValue,
};
pub type TransactionBuilder<'tx, Schema, Builder, State> =
crate::transaction::sqlite::typestate::TransactionBuilder<
'tx,
Transaction<Schema>,
Schema,
Builder,
State,
>;
pub struct Transaction<Schema = ()> {
tx: libsql::Transaction,
tx_type: SQLiteTransactionType,
savepoint_depth: AtomicU32,
schema: Schema,
}
impl<Schema> std::fmt::Debug for Transaction<Schema> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Transaction")
.field("tx_type", &self.tx_type)
.field("savepoint_depth", &self.savepoint_depth)
.finish_non_exhaustive()
}
}
impl<Schema> Transaction<Schema> {
pub(crate) const fn new(
tx: libsql::Transaction,
tx_type: SQLiteTransactionType,
schema: Schema,
) -> Self {
Self {
tx,
tx_type,
savepoint_depth: AtomicU32::new(0),
schema,
}
}
#[inline]
pub const fn schema(&self) -> &Schema {
&self.schema
}
#[inline]
pub const fn inner(&self) -> &libsql::Transaction {
&self.tx
}
#[inline]
pub const fn tx_type(&self) -> SQLiteTransactionType {
self.tx_type
}
async fn execute_raw(&self, sql: &str) -> Result<(), DrizzleError> {
self.tx.execute(sql, ()).await?;
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.savepoint_depth,
|sql| async move { self.execute_raw(&sql).await },
f(self),
)
.await
}
sqlite_transaction_constructors!();
pub async fn execute<'q, T>(&self, query: T) -> Result<u64, DrizzleError>
where
T: ToSQL<'q, SQLiteValue<'q>>,
{
let query = query.to_sql();
let (sql, params) = query.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
Ok(self.tx.execute(&sql, params).await?)
}
pub async fn all<'q, T, R>(&self, query: T) -> drizzle_core::error::Result<Vec<R>>
where
R: for<'r> TryFrom<&'r Row>,
for<'r> <R as TryFrom<&'r Row>>::Error: Into<DrizzleError>,
T: ToSQL<'q, SQLiteValue<'q>>,
{
self.rows(query).await?.collect().await
}
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<DrizzleError>,
T: ToSQL<'q, SQLiteValue<'q>>,
{
let sql = query.to_sql();
let (sql_str, params) = sql.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
let rows = self.tx.query(&sql_str, params).await?;
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<DrizzleError>,
T: ToSQL<'q, SQLiteValue<'q>>,
{
let sql = query.to_sql();
let (sql_str, params) = sql.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
let mut rows = self.tx.query(&sql_str, params).await?;
rows.next().await?.map_or_else(
|| Err(DrizzleError::NotFound),
|row| R::try_from(&row).map_err(Into::into),
)
}
pub async fn commit(self) -> Result<(), DrizzleError> {
Ok(self.tx.commit().await?)
}
pub async fn rollback(self) -> Result<(), DrizzleError> {
Ok(self.tx.rollback().await?)
}
}
#[cfg(feature = "libsql")]
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> {
let (sql, params) = self.builder.sql.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
Ok(self.runner.tx.execute(&sql, params).await?)
}
pub async fn all<R, Proof, AggProof>(self) -> drizzle_core::error::Result<Vec<R>>
where
for<'r> Mk: drizzle_core::row::DecodeSelectedRef<&'r ::libsql::Row, R>
+ drizzle_core::row::MarkerScopeValidFor<Proof>
+ drizzle_core::row::StrictDecodeMarker
+ drizzle_core::row::MarkerColumnCountValid<::libsql::Row, Rw, R>,
Mk: drizzle_core::row::MarkerAggValidFor<Grouped, AggProof>,
{
let (sql_str, params) = self.builder.sql.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
let mut rows = self.runner.tx.query(&sql_str, params).await?;
let mut decoded = Vec::new();
while let Some(row) = rows.next().await? {
decoded.push(<Mk as drizzle_core::row::DecodeSelectedRef<
&::libsql::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<DrizzleError>,
{
let (sql_str, params) = self.builder.sql.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
let rows = self.runner.tx.query(&sql_str, params).await?;
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 ::libsql::Row, R>
+ drizzle_core::row::MarkerScopeValidFor<Proof>
+ drizzle_core::row::StrictDecodeMarker
+ drizzle_core::row::MarkerColumnCountValid<::libsql::Row, Rw, R>,
Mk: drizzle_core::row::MarkerAggValidFor<Grouped, AggProof>,
{
let (sql_str, params) = self.builder.sql.build();
let params: Vec<libsql::Value> = params.into_iter().map(std::convert::Into::into).collect();
let mut rows = self.runner.tx.query(&sql_str, params).await?;
rows.next().await?.map_or_else(
|| Err(DrizzleError::NotFound),
|row| <Mk as drizzle_core::row::DecodeSelectedRef<&::libsql::Row, R>>::decode(&row),
)
}
}