use core::fmt;
use sqlx::migrate::Migrator;
use sqlx::postgres::PgConnection;
use sqlx::{Connection, Executor, PgPool};
static MIGRATOR: Migrator = sqlx::migrate!("./migrations");
#[derive(Clone, Copy, Debug)]
#[non_exhaustive]
pub struct MigrateOptions<'a> {
pub schema: &'a str,
}
impl Default for MigrateOptions<'_> {
fn default() -> Self {
Self { schema: "reliar" }
}
}
impl<'a> MigrateOptions<'a> {
#[must_use]
pub const fn schema(mut self, schema: &'a str) -> Self {
self.schema = schema;
self
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum MigrateError {
InvalidSchema {
schema: String,
},
Sqlx {
source: sqlx::migrate::MigrateError,
},
}
impl fmt::Display for MigrateError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InvalidSchema { schema } => write!(
f,
"{schema:?} is not a valid PostgreSQL identifier (expected \
[A-Za-z_][A-Za-z0-9_$]*, at most 63 bytes)"
),
Self::Sqlx { source } => write!(f, "migration failed: {source}"),
}
}
}
impl std::error::Error for MigrateError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Sqlx { source } => Some(source),
Self::InvalidSchema { .. } => None,
}
}
}
impl From<sqlx::migrate::MigrateError> for MigrateError {
fn from(source: sqlx::migrate::MigrateError) -> Self {
Self::Sqlx { source }
}
}
impl From<sqlx::Error> for MigrateError {
fn from(source: sqlx::Error) -> Self {
Self::Sqlx {
source: sqlx::migrate::MigrateError::Execute(source),
}
}
}
pub async fn migrate(pool: &PgPool, options: MigrateOptions<'_>) -> Result<(), MigrateError> {
if !crate::error::is_valid_schema_name(options.schema) {
return Err(MigrateError::InvalidSchema {
schema: options.schema.to_owned(),
});
}
let mut migrator = Migrator {
migrations: MIGRATOR.migrations.clone(),
ignore_missing: MIGRATOR.ignore_missing,
locking: MIGRATOR.locking,
no_tx: MIGRATOR.no_tx,
table_name: MIGRATOR.table_name.clone(),
create_schemas: MIGRATOR.create_schemas.clone(),
};
migrator.create_schema(options.schema.to_owned());
migrator.dangerous_set_table_name(format!("{}._migrations", options.schema));
migrator.set_locking(true);
let connect_options = pool.connect_options();
let mut conn = PgConnection::connect_with(&connect_options).await?;
conn.execute(sqlx::query(sqlx::AssertSqlSafe(format!(
"SET search_path = \"{}\", public",
options.schema.replace('"', "\"\"")
))))
.await?;
migrator.run(&mut conn).await?;
conn.close().await?;
Ok(())
}