pub mod connection;
pub mod conversion;
mod ext;
mod introspection;
pub mod transaction;
use async_trait::async_trait;
use sqlx::Executor;
use sqlx::postgres::PgPool;
use std::sync::Arc;
use super::provider::DatabaseProvider;
use crate::error::DatabaseResult;
use crate::models::{
DatabaseInfo, DatabaseTransaction, JsonRow, QueryResult, QuerySelector, ToDbValue,
};
use conversion::{bind_params, row_to_json, rows_to_result};
use transaction::PostgresTransaction;
#[derive(Debug)]
pub struct PostgresProvider {
pool: Arc<PgPool>,
}
impl PostgresProvider {
pub async fn new(database_url: &str) -> DatabaseResult<Self> {
Self::new_with_pool(database_url, &connection::PoolConfig::default()).await
}
pub async fn new_with_pool(
database_url: &str,
pool_config: &connection::PoolConfig,
) -> DatabaseResult<Self> {
let connect_options = connection::connect_options(database_url)?;
let pool = connection::connect_with_retry(
connection::build_pool_options(pool_config),
connect_options,
)
.await?;
Ok(Self {
pool: Arc::new(pool),
})
}
#[must_use]
pub const fn from_pool(pool: Arc<PgPool>) -> Self {
Self { pool }
}
#[must_use]
pub fn pool(&self) -> &PgPool {
&self.pool
}
}
#[async_trait]
impl DatabaseProvider for PostgresProvider {
fn get_postgres_pool(&self) -> Arc<PgPool> {
Arc::clone(&self.pool)
}
async fn execute(
&self,
query: &dyn QuerySelector,
params: &[&dyn ToDbValue],
) -> DatabaseResult<u64> {
let sql = query.select_query();
let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
let query_obj = bind_params(query_obj, params);
let result = query_obj.execute(&*self.pool).await?;
Ok(result.rows_affected())
}
async fn execute_raw(&self, sql: &str) -> DatabaseResult<()> {
let mut conn = self.pool.acquire().await?;
conn.execute(sqlx::AssertSqlSafe(sql.to_owned())).await?;
Ok(())
}
async fn fetch_all(
&self,
query: &dyn QuerySelector,
params: &[&dyn ToDbValue],
) -> DatabaseResult<Vec<JsonRow>> {
let sql = query.select_query();
let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
let query_obj = bind_params(query_obj, params);
let rows = query_obj.fetch_all(&*self.pool).await?;
Ok(rows.iter().map(row_to_json).collect())
}
async fn fetch_one(
&self,
query: &dyn QuerySelector,
params: &[&dyn ToDbValue],
) -> DatabaseResult<JsonRow> {
let sql = query.select_query();
let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
let query_obj = bind_params(query_obj, params);
let row = query_obj.fetch_one(&*self.pool).await?;
Ok(row_to_json(&row))
}
async fn fetch_optional(
&self,
query: &dyn QuerySelector,
params: &[&dyn ToDbValue],
) -> DatabaseResult<Option<JsonRow>> {
let sql = query.select_query();
let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
let query_obj = bind_params(query_obj, params);
let row = query_obj.fetch_optional(&*self.pool).await?;
Ok(row.map(|r| row_to_json(&r)))
}
async fn begin_transaction(&self) -> DatabaseResult<Box<dyn DatabaseTransaction>> {
let tx = self.pool.begin().await?;
Ok(Box::new(PostgresTransaction::new(tx)))
}
async fn get_database_info(&self) -> DatabaseResult<DatabaseInfo> {
introspection::get_database_info(&self.pool).await
}
async fn test_connection(&self) -> DatabaseResult<()> {
sqlx::query("SELECT 1").fetch_one(&*self.pool).await?;
Ok(())
}
async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
let statements = crate::services::SqlExecutor::parse_sql_statements(sql)?;
for statement in statements {
sqlx::query(sqlx::AssertSqlSafe(statement))
.execute(&*self.pool)
.await?;
}
Ok(())
}
async fn query_raw(&self, query: &dyn QuerySelector) -> DatabaseResult<QueryResult> {
let sql = query.select_query();
let start = std::time::Instant::now();
let rows = sqlx::query(sqlx::AssertSqlSafe(sql))
.fetch_all(&*self.pool)
.await?;
Ok(rows_to_result(rows, start))
}
async fn query_raw_with(
&self,
query: &dyn QuerySelector,
params: &[&dyn ToDbValue],
) -> DatabaseResult<QueryResult> {
let sql = query.select_query();
let start = std::time::Instant::now();
let query_obj = bind_params(sqlx::query(sqlx::AssertSqlSafe(sql)), params);
let rows = query_obj.fetch_all(&*self.pool).await?;
Ok(rows_to_result(rows, start))
}
}