pub mod dsn;
pub mod testcontainers;
pub mod validate;
use std::future::Future;
use std::pin::Pin;
pub type PgMajor = u32;
type ResetFuture<'a> = Pin<Box<dyn Future<Output = anyhow::Result<()>> + Send + 'a>>;
type CheckoutFuture<'a> =
Pin<Box<dyn Future<Output = anyhow::Result<Box<dyn ShadowGuard>>> + Send + 'a>>;
pub trait ShadowGuard: Send {
fn url(&self) -> &str;
fn reset(&mut self) -> ResetFuture<'_>;
}
pub trait ShadowBackend: Send + Sync {
fn checkout(&self, major: PgMajor) -> CheckoutFuture<'_>;
}
pub fn resolve(config: &crate::config::ShadowConfig) -> anyhow::Result<Box<dyn ShadowBackend>> {
match config.backend.as_deref().unwrap_or("auto") {
"testcontainers" => Ok(Box::new(testcontainers::TestcontainersBackend::new(config))),
"dsn" => Ok(Box::new(dsn::DsnBackend::new(config)?)),
"auto" => {
if config.url.is_some() || config.url_env.is_some() {
Ok(Box::new(dsn::DsnBackend::new(config)?))
} else if docker_available() {
Ok(Box::new(testcontainers::TestcontainersBackend::new(config)))
} else {
anyhow::bail!(
"no shadow backend available: configure [shadow].url or install Docker"
)
}
}
other => anyhow::bail!("unknown shadow backend: {other}"),
}
}
#[must_use]
pub fn docker_available() -> bool {
if std::env::var_os("PGEVOLVE_DISABLE_DOCKER_TESTS").is_some() {
return false;
}
std::process::Command::new("docker")
.arg("info")
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.status()
.is_ok_and(|s| s.success())
}
pub(super) async fn install_extensions(url: &str, extensions: &[String]) -> anyhow::Result<()> {
for ext in extensions {
validate_extension_name(ext)?;
}
let (client, conn) = tokio_postgres::connect(url, tokio_postgres::NoTls).await?;
tokio::spawn(conn);
for ext in extensions {
let stmt = format!(
"CREATE EXTENSION IF NOT EXISTS \"{}\"",
ext.replace('"', "\"\"")
);
client.batch_execute(&stmt).await?;
}
Ok(())
}
fn validate_extension_name(name: &str) -> anyhow::Result<()> {
if name.is_empty() {
anyhow::bail!("[shadow].extensions: empty extension name");
}
let valid = name.chars().enumerate().all(|(i, c)| {
if i == 0 {
c.is_ascii_alphabetic() || c == '_'
} else {
c.is_ascii_alphanumeric() || c == '_' || c == '-'
}
});
if !valid {
anyhow::bail!(
"[shadow].extensions: {name:?} is not a valid extension identifier \
(expected [a-zA-Z_][a-zA-Z0-9_-]*)",
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_accepts_valid_names() {
for valid in &["pg_trgm", "uuid-ossp", "vector", "_underscore_first", "p"] {
assert!(
validate_extension_name(valid).is_ok(),
"{valid} should validate"
);
}
}
#[test]
fn validate_rejects_invalid_names() {
for bad in &[
"",
"1leading_digit",
"-leading_hyphen",
"has space",
"semicolon;here",
"pg_trgm; DROP TABLE x",
] {
assert!(
validate_extension_name(bad).is_err(),
"{bad:?} should reject"
);
}
}
}