use std::sync::{Mutex, OnceLock};
use color_eyre::eyre::WrapErr;
use dashmap::DashMap;
use tracing::info_span;
use super::connection::{TestClusterConnection, escape_identifier};
use super::temporary_database::TemporaryDatabase;
use crate::error::BootstrapResult;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct DatabaseName(String);
impl DatabaseName {
#[must_use]
pub fn new(name: impl Into<String>) -> Self {
Self(name.into())
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
impl AsRef<str> for DatabaseName {
fn as_ref(&self) -> &str {
&self.0
}
}
impl From<&str> for DatabaseName {
fn from(s: &str) -> Self {
Self(s.to_owned())
}
}
impl From<String> for DatabaseName {
fn from(s: String) -> Self {
Self(s)
}
}
static TEMPLATE_LOCKS: OnceLock<DashMap<String, Mutex<()>>> = OnceLock::new();
fn template_locks() -> &'static DashMap<String, Mutex<()>> {
TEMPLATE_LOCKS.get_or_init(DashMap::new)
}
impl TestClusterConnection {
fn execute_ddl_command(
&self,
sql_template: &str,
name: &str,
error_msg_verb: &str,
) -> BootstrapResult<()> {
let mut client = self.admin_client()?;
let escaped = escape_identifier(name);
let sql = sql_template.replace("{}", &format!("\"{escaped}\""));
client
.batch_execute(&sql)
.wrap_err(format!("failed to {error_msg_verb} database '{name}'"))
.map_err(crate::error::BootstrapError::from)
}
pub fn create_database(&self, name: impl Into<DatabaseName>) -> BootstrapResult<()> {
let db_name = name.into();
let _span = info_span!("create_database", db = %db_name.as_str()).entered();
self.execute_ddl_command("CREATE DATABASE {}", db_name.as_str(), "create")
}
pub fn create_database_from_template(
&self,
name: impl Into<DatabaseName>,
template: impl Into<DatabaseName>,
) -> BootstrapResult<()> {
let db_name = name.into();
let template_name = template.into();
let _span =
info_span!("create_database_from_template", db = %db_name.as_str(), template = %template_name.as_str()).entered();
let mut client = self.admin_client()?;
let escaped_name = escape_identifier(db_name.as_str());
let escaped_template = escape_identifier(template_name.as_str());
let sql = format!("CREATE DATABASE \"{escaped_name}\" TEMPLATE \"{escaped_template}\"");
client
.batch_execute(&sql)
.wrap_err(format!(
"failed to create database '{}' from template '{}'",
db_name.as_str(),
template_name.as_str()
))
.map_err(crate::error::BootstrapError::from)
}
pub fn drop_database(&self, name: impl Into<DatabaseName>) -> BootstrapResult<()> {
let db_name = name.into();
let _span = info_span!("drop_database", db = %db_name.as_str()).entered();
self.execute_ddl_command("DROP DATABASE {}", db_name.as_str(), "drop")
}
pub fn database_exists(&self, name: impl Into<DatabaseName>) -> BootstrapResult<bool> {
let db_name = name.into();
let mut client = self.admin_client()?;
let row = client
.query_one(
"SELECT EXISTS(SELECT 1 FROM pg_database WHERE datname = $1)",
&[&db_name.as_str()],
)
.wrap_err("failed to query pg_database")
.map_err(crate::error::BootstrapError::from)?;
Ok(row.get(0))
}
pub fn ensure_template_exists<F>(
&self,
name: impl Into<DatabaseName>,
setup_fn: F,
) -> BootstrapResult<()>
where
F: FnOnce(&str) -> BootstrapResult<()>,
{
let db_name = name.into();
let _span = info_span!("ensure_template_exists", template = %db_name.as_str()).entered();
let locks = template_locks();
let lock = locks
.entry(db_name.as_str().to_owned())
.or_insert_with(|| Mutex::new(()));
let _guard = lock
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !self.database_exists(db_name.as_str())? {
self.create_database(db_name.as_str())?;
setup_fn(db_name.as_str())?;
}
Ok(())
}
pub fn temporary_database(
&self,
name: impl Into<DatabaseName>,
) -> BootstrapResult<TemporaryDatabase> {
let db_name = name.into();
let _span = info_span!("temporary_database", db = %db_name.as_str()).entered();
self.create_database(db_name.as_str())?;
Ok(TemporaryDatabase::new(
db_name.as_str().to_owned(),
self.database_url("postgres"),
self.database_url(db_name.as_str()),
))
}
pub fn temporary_database_from_template(
&self,
name: impl Into<DatabaseName>,
template: impl Into<DatabaseName>,
) -> BootstrapResult<TemporaryDatabase> {
let db_name = name.into();
let template_name = template.into();
let _span =
info_span!("temporary_database_from_template", db = %db_name.as_str(), template = %template_name.as_str())
.entered();
self.create_database_from_template(db_name.as_str(), template_name.as_str())?;
Ok(TemporaryDatabase::new(
db_name.as_str().to_owned(),
self.database_url("postgres"),
self.database_url(db_name.as_str()),
))
}
}