use std::fmt;
use std::sync::Arc;
use sqlx::pool::PoolConnection;
use sqlx::postgres::{PgConnectOptions, PgPool, PgPoolOptions};
use sqlx::{Executor, Postgres, Transaction};
use turnframe_store::error::StoreError;
use turnframe_store::stores::{Stores, StoresBuilderError};
use crate::config::PgStoreConfig;
use crate::error::{commit_failed, store_error};
pub static MIGRATOR: sqlx::migrate::Migrator = sqlx::migrate!("./migrations");
#[derive(Clone)]
pub struct PgStores {
pool: PgPool,
schema: Option<String>,
}
impl PgStores {
pub async fn connect(url: &str) -> Result<Self, StoreError> {
Self::connect_with(url, &PgStoreConfig::new()).await
}
pub async fn connect_with(url: &str, config: &PgStoreConfig) -> Result<Self, StoreError> {
let options: PgConnectOptions = url.parse().map_err(|error| store_error(&error))?;
let pool = config
.apply_pool(PgPoolOptions::new())
.connect_with(config.apply_connection(options))
.await
.map_err(|error| store_error(&error))?;
Ok(Self {
pool,
schema: config.schema_name().map(ToOwned::to_owned),
})
}
#[must_use]
pub fn from_pool(pool: PgPool) -> Self {
Self { pool, schema: None }
}
pub fn with_schema(mut self, schema: impl Into<String>) -> Result<Self, crate::ConfigError> {
let config = PgStoreConfig::new().schema(schema)?;
self.schema = config.schema_name().map(ToOwned::to_owned);
Ok(self)
}
pub async fn migrate(&self) -> Result<(), sqlx::migrate::MigrateError> {
if let Some(schema) = &self.schema {
let statement = format!("CREATE SCHEMA IF NOT EXISTS {schema}");
match self.pool.execute(statement.as_str()).await {
Ok(_) => {}
Err(error) if created_concurrently(&error) => {}
Err(error) => return Err(error.into()),
}
}
tracing::debug!(
migrations = MIGRATOR.iter().len(),
"applying turnframe store migrations"
);
MIGRATOR.run(&self.pool).await
}
#[must_use]
pub fn pool(&self) -> &PgPool {
&self.pool
}
#[must_use]
pub fn schema(&self) -> Option<&str> {
self.schema.as_deref()
}
pub async fn close(&self) {
self.pool.close().await;
}
pub fn stores(&self) -> Result<Stores, StoresBuilderError> {
let backend = Arc::new(self.clone());
Stores::builder()
.conversations(backend.clone())
.interactions(backend.clone())
.journal(backend.clone())
.events(backend.clone())
.outbox(backend.clone())
.replay(backend.clone())
.commit(backend)
.build()
}
pub(crate) async fn connection(&self) -> Result<PoolConnection<Postgres>, StoreError> {
self.pool
.acquire()
.await
.map_err(|error| store_error(&error))
}
pub(crate) async fn transaction(&self) -> Result<Transaction<'_, Postgres>, StoreError> {
self.pool.begin().await.map_err(|error| store_error(&error))
}
}
fn created_concurrently(error: &sqlx::Error) -> bool {
const DUPLICATE_SCHEMA: &str = "42P06";
const UNIQUE_VIOLATION: &str = "23505";
error
.as_database_error()
.and_then(|database| database.code())
.is_some_and(|code| code == DUPLICATE_SCHEMA || code == UNIQUE_VIOLATION)
}
pub(crate) async fn commit(transaction: Transaction<'_, Postgres>) -> Result<(), StoreError> {
transaction
.commit()
.await
.map_err(|error| commit_failed(&error))
}
impl fmt::Debug for PgStores {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PgStores")
.field("schema", &self.schema.as_deref().unwrap_or("<default>"))
.field("connections", &self.pool.size())
.field("closed", &self.pool.is_closed())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_migrations_are_embedded() {
assert_eq!(MIGRATOR.iter().len(), 2, "both migrations are embedded");
assert!(
MIGRATOR
.iter()
.all(|migration| !migration.sql.trim().is_empty())
);
}
#[tokio::test]
async fn a_schema_name_is_validated_before_it_reaches_ddl() {
let store = PgStores::from_pool(PgPool::connect_lazy("postgres://x/y").unwrap());
assert!(store.clone().with_schema("tf_test").is_ok());
assert!(store.with_schema("tf\";DROP SCHEMA public").is_err());
}
#[tokio::test]
async fn debug_output_carries_no_connection_string() {
let store =
PgStores::from_pool(PgPool::connect_lazy("postgres://secret:hunter2@host/db").unwrap());
let rendered = format!("{store:?}");
assert!(!rendered.contains("hunter2"), "{rendered}");
assert!(rendered.contains("<default>"), "{rendered}");
}
}