1use crate::handle::DbHandle;
2use crate::DbError;
3use sea_orm::{Database, DatabaseConnection, DbErr, TransactionTrait};
4use std::sync::Arc;
5
6pub struct TestDb {
8 _conn: DatabaseConnection,
9 handle: DbHandle,
10 tx: Option<Arc<sea_orm::DatabaseTransaction>>,
11}
12
13impl TestDb {
14 pub fn db(&self) -> &DbHandle {
15 &self.handle
16 }
17
18 pub async fn rollback(mut self) -> Result<(), DbError> {
19 if let Some(arc) = self.tx.take() {
20 match Arc::try_unwrap(arc) {
21 Ok(tx) => tx.rollback().await.map_err(DbError)?,
22 Err(_) => {
23 return Err(DbError(DbErr::Custom(
24 "test transaction still held".into(),
25 )));
26 }
27 }
28 }
29 Ok(())
30 }
31}
32
33pub async fn test_db() -> Result<TestDb, DbError> {
35 let url = std::env::var("DATABASE_URL")
36 .map_err(|_| DbError(DbErr::Custom("DATABASE_URL is not set".into())))?;
37 let conn = Database::connect(&url).await.map_err(DbError)?;
38 let tx = conn.begin().await.map_err(DbError)?;
39 let arc = Arc::new(tx);
40 Ok(TestDb {
41 _conn: conn,
42 handle: DbHandle::Tx(Arc::clone(&arc)),
43 tx: Some(arc),
44 })
45}