#![allow(dead_code)]
use uuid::Uuid;
use arcature_db::{Db, DbConfig};
pub struct TestDb {
db: Db,
maintenance: Option<Maintenance>,
}
struct Maintenance {
db: Db,
db_name: String,
}
pub enum DbOrSkip {
Db(TestDb),
Skipped,
}
pub async fn require_db() -> DbOrSkip {
let Ok(maintenance_url) = std::env::var("ARCATURE_TEST_DB_URL") else {
eprintln!("skipping: ARCATURE_TEST_DB_URL not set");
return DbOrSkip::Skipped;
};
if !is_localhost(&maintenance_url) {
eprintln!("skipping: ARCATURE_TEST_DB_URL is not localhost");
return DbOrSkip::Skipped;
}
let db_name = format!("arcature_data_test_{}", Uuid::new_v4().simple());
let maintenance_config = match DbConfig::new(&maintenance_url) {
Ok(config) => config,
Err(error) => {
eprintln!("skipping: cannot parse maintenance URL: {error}");
return DbOrSkip::Skipped;
}
};
let maintenance = match Db::connect(maintenance_config).await {
Ok(db) => db,
Err(error) => {
eprintln!("skipping: cannot connect maintenance pool: {error}");
return DbOrSkip::Skipped;
}
};
let create = arcature_db::sqlx::query(arcature_db::sqlx::AssertSqlSafe(format!(
"CREATE DATABASE \"{db_name}\""
)))
.execute(maintenance.sqlx())
.await;
if let Err(error) = create {
eprintln!("skipping: cannot create test database: {error}");
maintenance.close().await;
return DbOrSkip::Skipped;
}
let test_url = replace_database(&maintenance_url, &db_name);
let test_config = match DbConfig::new(&test_url) {
Ok(config) => config,
Err(error) => {
eprintln!("skipping: cannot parse test URL: {error}");
let _ = drop_test_db(&maintenance, &db_name).await;
maintenance.close().await;
return DbOrSkip::Skipped;
}
};
let db = match Db::connect(test_config).await {
Ok(db) => db,
Err(error) => {
eprintln!("skipping: cannot connect test pool: {error}");
let _ = drop_test_db(&maintenance, &db_name).await;
maintenance.close().await;
return DbOrSkip::Skipped;
}
};
DbOrSkip::Db(TestDb {
db,
maintenance: Some(Maintenance {
db: maintenance,
db_name,
}),
})
}
impl TestDb {
#[must_use]
pub fn db(&self) -> &Db {
&self.db
}
pub async fn stop(mut self) {
self.db.close().await;
if let Some(maintenance) = self.maintenance.take() {
let _ = drop_test_db(&maintenance.db, &maintenance.db_name).await;
maintenance.db.close().await;
}
}
}
async fn drop_test_db(maintenance: &Db, db_name: &str) -> Result<(), arcature_db::sqlx::Error> {
arcature_db::sqlx::query(
"SELECT pg_terminate_backend(pid) FROM pg_stat_activity \
WHERE datname = $1 AND pid <> pg_backend_pid()",
)
.bind(db_name)
.execute(maintenance.sqlx())
.await?;
arcature_db::sqlx::query(arcature_db::sqlx::AssertSqlSafe(format!(
"DROP DATABASE IF EXISTS \"{db_name}\""
)))
.execute(maintenance.sqlx())
.await?;
Ok(())
}
pub async fn setup_schema(db: &Db) -> Result<(), arcature_db::sqlx::Error> {
arcature_db::sqlx::query(
"CREATE TABLE IF NOT EXISTS users (\
id SERIAL PRIMARY KEY,\
email TEXT NOT NULL UNIQUE,\
name TEXT NOT NULL,\
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()\
)",
)
.execute(db.sqlx())
.await?;
arcature_db::sqlx::query(
"CREATE TABLE IF NOT EXISTS posts (\
id SERIAL PRIMARY KEY,\
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,\
title TEXT NOT NULL,\
body TEXT,\
active BOOLEAN NOT NULL DEFAULT TRUE,\
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()\
)",
)
.execute(db.sqlx())
.await?;
Ok(())
}
pub async fn reset_rows(db: &Db) -> Result<(), arcature_db::sqlx::Error> {
arcature_db::sqlx::query("TRUNCATE TABLE posts, users RESTART IDENTITY CASCADE")
.execute(db.sqlx())
.await?;
Ok(())
}
fn is_localhost(url: &str) -> bool {
let after_scheme = url.split("://").nth(1).unwrap_or(url);
let authority_end = after_scheme
.find(['/', '?', '#'])
.unwrap_or(after_scheme.len());
let authority = &after_scheme[..authority_end];
let host = authority.rsplit('@').next().unwrap_or(authority);
let host = host.rsplit_once(':').map(|(h, _)| h).unwrap_or(host);
host == "localhost" || host == "127.0.0.1" || host == "::1" || host.ends_with(".localhost")
}
fn replace_database(url: &str, db_name: &str) -> String {
let scheme_end = url.find("://").map(|i| i + 3).unwrap_or(0);
let path_start = url[scheme_end..]
.find('/')
.map(|i| i + scheme_end)
.unwrap_or(url.len());
let query_start = url[path_start..].find('?').map(|i| i + path_start);
let base = &url[..path_start];
let query = match query_start {
Some(qs) => &url[qs..],
None => "",
};
format!("{base}/{db_name}{query}")
}
#[cfg(test)]
mod tests {
use super::replace_database;
#[test]
fn replace_database_basic() {
assert_eq!(
replace_database("postgres://u:p@localhost:5432/main", "newdb"),
"postgres://u:p@localhost:5432/newdb"
);
}
#[test]
fn replace_database_preserves_query() {
assert_eq!(
replace_database("postgres://localhost/main?sslmode=disable", "newdb"),
"postgres://localhost/newdb?sslmode=disable"
);
}
#[test]
fn replace_database_when_no_path() {
assert_eq!(
replace_database("postgres://localhost", "newdb"),
"postgres://localhost/newdb"
);
}
}