use std::future::Future;
use std::pin::Pin;
use async_trait::async_trait;
use turso_sql::Statement;
use crate::database::Database;
use crate::error::Result;
use crate::executor::{self, Conn, ExecResult, Row, RowStream};
use crate::transaction::{Transaction, TransactionMode};
#[async_trait]
pub trait ConnectionTrait: Send + Sync {
async fn execute(&self, statement: Statement) -> Result<ExecResult>;
async fn execute_unprepared(&self, sql: &str) -> Result<ExecResult>;
async fn query_one(&self, statement: Statement) -> Result<Option<Row>>;
async fn query_all(&self, statement: Statement) -> Result<Vec<Row>>;
}
pub trait StreamTrait: Send + Sync {
fn stream<'a>(
&'a self,
statement: Statement,
) -> Pin<Box<dyn Future<Output = Result<RowStream<'a>>> + Send + 'a>>;
}
#[async_trait]
pub trait TransactionTrait: Send + Sync {
async fn begin(&self) -> Result<Transaction>;
async fn begin_with_mode(&self, mode: TransactionMode) -> Result<Transaction>;
async fn transaction<F, T, E>(&self, callback: F) -> std::result::Result<T, E>
where
F: for<'c> FnOnce(
&'c Transaction,
)
-> Pin<Box<dyn Future<Output = std::result::Result<T, E>> + Send + 'c>>
+ Send,
T: Send,
E: From<crate::Error> + Send,
{
let txn = self.begin().await?;
match callback(&txn).await {
Ok(value) => {
txn.commit().await?;
Ok(value)
}
Err(err) => {
txn.rollback().await?;
Err(err)
}
}
}
}
#[async_trait]
impl ConnectionTrait for Database {
async fn execute(&self, statement: Statement) -> Result<ExecResult> {
let conn = self.acquire().await?;
executor::execute(&conn, &statement).await
}
async fn execute_unprepared(&self, sql: &str) -> Result<ExecResult> {
let conn = self.acquire().await?;
executor::execute_unprepared(&conn, sql).await
}
async fn query_one(&self, statement: Statement) -> Result<Option<Row>> {
let conn = self.acquire().await?;
executor::query_one(&conn, &statement).await
}
async fn query_all(&self, statement: Statement) -> Result<Vec<Row>> {
let conn = self.acquire().await?;
executor::query_all(&conn, &statement).await
}
}
impl StreamTrait for Database {
fn stream<'a>(
&'a self,
statement: Statement,
) -> Pin<Box<dyn Future<Output = Result<RowStream<'a>>> + Send + 'a>> {
Box::pin(async move {
let conn = self.acquire().await?;
let raw: Conn = (*conn).clone();
executor::stream(&raw, &statement, conn).await
})
}
}
#[async_trait]
impl TransactionTrait for Database {
async fn begin(&self) -> Result<Transaction> {
self.begin_with_mode(TransactionMode::Deferred).await
}
async fn begin_with_mode(&self, mode: TransactionMode) -> Result<Transaction> {
let conn = self.acquire().await?;
Transaction::begin_top(conn, mode).await
}
}
macro_rules! forward_ref {
($($t:ty),*) => {$(
#[async_trait]
impl ConnectionTrait for &$t {
async fn execute(&self, statement: Statement) -> Result<ExecResult> {
(**self).execute(statement).await
}
async fn execute_unprepared(&self, sql: &str) -> Result<ExecResult> {
(**self).execute_unprepared(sql).await
}
async fn query_one(&self, statement: Statement) -> Result<Option<Row>> {
(**self).query_one(statement).await
}
async fn query_all(&self, statement: Statement) -> Result<Vec<Row>> {
(**self).query_all(statement).await
}
}
)*};
}
forward_ref!(Database, Transaction);