use async_trait::async_trait;
use std::fmt::Debug;
use std::marker::PhantomData;
use crate::testdb::{
DatabaseBackend, DatabasePool, TestDatabaseConnection,
transaction::{DBTransactionManager, DatabaseTransaction},
};
use crate::TestContext;
pub trait TransactionStarter<DB: DatabaseBackend> {
type Transaction: DatabaseTransaction<Error = DB::Error> + Send + Sync + 'static;
type Connection: TestDatabaseConnection + Send + Sync + 'static;
fn begin_transaction_type() -> Self::Transaction;
}
impl<DB> TransactionStarter<DB> for TestContext<DB>
where
DB: DatabaseBackend + Send + Sync + Debug + 'static,
{
type Transaction = MockTransactionFor<DB>;
type Connection = MockConnectionFor<DB>;
fn begin_transaction_type() -> Self::Transaction {
panic!("This method should never be called")
}
}
pub struct MockTransactionFor<DB: DatabaseBackend>(PhantomData<DB>);
pub struct MockConnectionFor<DB: DatabaseBackend>(PhantomData<DB>);
#[async_trait]
impl<DB: DatabaseBackend> DatabaseTransaction for MockTransactionFor<DB> {
type Error = DB::Error;
async fn commit(&mut self) -> Result<(), Self::Error> {
Ok(())
}
async fn rollback(&mut self) -> Result<(), Self::Error> {
Ok(())
}
}
impl<DB: DatabaseBackend> TestDatabaseConnection for MockConnectionFor<DB> {
fn connection_string(&self) -> String {
"mock".to_string()
}
}
#[async_trait]
impl<DB, T, Conn> DBTransactionManager<T, Conn> for TestContext<DB>
where
DB: DatabaseBackend + Send + Sync + Debug + 'static,
T: DatabaseTransaction<Error = DB::Error> + Send + Sync + 'static,
Conn: TestDatabaseConnection + Send + Sync + 'static,
DB::Pool: DatabasePool<Connection = Conn, Error = DB::Error>,
{
type Error = DB::Error;
type Tx = T;
async fn begin_transaction(&mut self) -> Result<Self::Tx, Self::Error> {
Err(From::from(
"Transaction implementation is database-specific and must be provided for each backend"
.to_string(),
))
}
async fn commit_transaction(tx: &mut Self::Tx) -> Result<(), Self::Error> {
tx.commit().await
}
async fn rollback_transaction(tx: &mut Self::Tx) -> Result<(), Self::Error> {
tx.rollback().await
}
}