use std::{
env,
time::{SystemTime, UNIX_EPOCH}
};
use sqlx::{
AssertSqlSafe, Connection, PgConnection, PgPool,
postgres::{PgConnectOptions, PgPoolOptions}
};
use uuid::Uuid;
const URL_VARS: [&str; 2] = ["ENTITY_DERIVE_TEST_DATABASE_URL", "DATABASE_URL"];
const DB_PREFIX: &str = "ed_t_";
const STALE_AFTER_MS: u128 = 30 * 60 * 1000;
const CLOSE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
pub struct TestDb {
name: String,
pool: PgPool,
admin: PgConnectOptions
}
impl TestDb {
pub const fn pool(&self) -> &PgPool {
&self.pool
}
pub async fn run(&self, sql: &str) {
sqlx::raw_sql(AssertSqlSafe(sql.to_owned()))
.execute(&self.pool)
.await
.unwrap_or_else(|e| panic!("failed to run SQL against {}: {e}\n{sql}", self.name));
}
pub async fn teardown(self) {
let _ = tokio::time::timeout(CLOSE_TIMEOUT, self.pool.close()).await;
let Ok(mut conn) = PgConnection::connect_with(&self.admin).await else {
return;
};
let drop_sql = format!("DROP DATABASE IF EXISTS \"{}\" WITH (FORCE)", self.name);
let _ = sqlx::raw_sql(AssertSqlSafe(drop_sql))
.execute(&mut conn)
.await;
}
}
pub async fn provision(label: &str, migrations: &[&str]) -> Option<TestDb> {
let base = maintenance_url()?;
let admin: PgConnectOptions = base
.parse()
.unwrap_or_else(|e| panic!("invalid Postgres URL in the environment: {e}"));
sweep_abandoned(&admin).await;
let name = database_name(label);
let mut conn = PgConnection::connect_with(&admin)
.await
.unwrap_or_else(|e| panic!("cannot reach the Postgres server: {e}"));
let create_sql = format!("CREATE DATABASE \"{name}\"");
sqlx::raw_sql(AssertSqlSafe(create_sql))
.execute(&mut conn)
.await
.unwrap_or_else(|e| panic!("cannot create database {name}: {e}"));
let _ = conn.close().await;
let pool = PgPoolOptions::new()
.max_connections(4)
.connect_with(admin.clone().database(&name))
.await
.unwrap_or_else(|e| panic!("cannot connect to {name}: {e}"));
let db = TestDb {
name,
pool,
admin
};
for script in migrations {
db.run(script).await;
}
Some(db)
}
fn maintenance_url() -> Option<String> {
for key in URL_VARS {
if let Ok(value) = env::var(key)
&& !value.trim().is_empty()
{
return Some(value);
}
}
assert!(
env::var("CI").is_err(),
"the live-Postgres suite requires {} under CI; the job must provide a server",
URL_VARS.join(" or ")
);
eprintln!(
"skipping: no live Postgres configured, set {} to run this test",
URL_VARS.join(" or ")
);
None
}
fn database_name(label: &str) -> String {
let millis = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| d.as_millis());
let slug: String = label
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() {
c.to_ascii_lowercase()
} else {
'_'
}
})
.take(16)
.collect();
let salt = Uuid::new_v4().simple().to_string();
format!("{DB_PREFIX}{millis}_{}_{slug}", &salt[..8])
}
async fn sweep_abandoned(admin: &PgConnectOptions) {
let Ok(mut conn) = PgConnection::connect_with(admin).await else {
return;
};
let listed: Result<Vec<(String,)>, _> =
sqlx::query_as("SELECT datname FROM pg_database WHERE datname LIKE $1")
.bind(format!("{DB_PREFIX}%"))
.fetch_all(&mut conn)
.await;
let Ok(rows) = listed else {
return;
};
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| d.as_millis());
for (name,) in rows {
let abandoned = created_at_millis(&name)
.is_some_and(|created| now.saturating_sub(created) > STALE_AFTER_MS);
if abandoned {
let drop_sql = format!("DROP DATABASE IF EXISTS \"{name}\" WITH (FORCE)");
let _ = sqlx::raw_sql(AssertSqlSafe(drop_sql))
.execute(&mut conn)
.await;
}
}
let _ = conn.close().await;
}
fn created_at_millis(name: &str) -> Option<u128> {
name.strip_prefix(DB_PREFIX)?
.split('_')
.next()?
.parse()
.ok()
}