use sqlx::any::AnyPoolOptions;
use crate::migration::{run, DbBackend, MigrationSet};
use crate::Db;
const ENV_URL: &str = "LATERITE_TEST_DATABASE_URL";
pub struct TestGuard {
inner: Option<Cleanup>,
}
struct Cleanup {
admin_url: String,
db_name: String,
backend: DbBackend,
}
impl Drop for TestGuard {
fn drop(&mut self) {
let Some(cleanup) = self.inner.take() else {
return;
};
let _ = std::thread::spawn(move || {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("teardown runtime");
rt.block_on(async move {
if let Ok(admin) = AnyPoolOptions::new()
.max_connections(1)
.connect(&cleanup.admin_url)
.await
{
if matches!(cleanup.backend, DbBackend::Postgres) {
let _ = sqlx::query(&format!(
"select pg_terminate_backend(pid) from pg_stat_activity \
where datname = '{}' and pid <> pg_backend_pid()",
cleanup.db_name
))
.execute(&admin)
.await;
}
let _ = sqlx::query(&format!("drop database if exists {}", cleanup.db_name))
.execute(&admin)
.await;
}
});
})
.join();
}
}
pub async fn connect_test(migrations: &[MigrationSet]) -> (Db, TestGuard) {
sqlx::any::install_default_drivers();
let base = std::env::var(ENV_URL).unwrap_or_else(|_| "sqlite::memory:".to_string());
let backend = DbBackend::from_url(&base).expect("recognised database URL");
let (db, guard) = match backend {
DbBackend::Sqlite => {
let pool = AnyPoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("connect in-memory sqlite");
(Db::new(pool, backend), TestGuard { inner: None })
}
DbBackend::Postgres | DbBackend::Mysql => {
let db_name = format!("laterite_test_{}", uuid::Uuid::new_v4().simple());
let admin = AnyPoolOptions::new()
.max_connections(1)
.connect(&base)
.await
.expect("connect maintenance database");
sqlx::query(&format!("create database {db_name}"))
.execute(&admin)
.await
.expect("create ephemeral test database");
admin.close().await;
let pool = AnyPoolOptions::new()
.max_connections(5)
.connect(&swap_database(&base, &db_name))
.await
.expect("connect ephemeral test database");
let guard = TestGuard {
inner: Some(Cleanup {
admin_url: base,
db_name,
backend,
}),
};
(Db::new(pool, backend), guard)
}
};
run(&db.pool, db.backend, migrations)
.await
.expect("apply migrations to the test database");
(db, guard)
}
fn swap_database(base: &str, name: &str) -> String {
let mut url = url::Url::parse(base).expect("valid database url");
url.set_path(name);
url.into()
}