use std::time::Duration;
use sqlx::mysql::MySqlPoolOptions;
use sqlx::postgres::PgPoolOptions;
use sqlx::sqlite::SqlitePoolOptions;
use super::lru_cache::LruCache;
use crate::connector::DbConnectorConfig;
use crate::errors::OrionError;
#[derive(Clone)]
pub enum SqlPool {
Postgres(sqlx::PgPool),
MySql(sqlx::MySqlPool),
Sqlite(sqlx::SqlitePool),
}
macro_rules! dispatch_sql_pool {
($self:expr, $p:ident, $rows_to_json:ident, $bind:ident, $typed_args:ident, $write_result:ident
=> $body:expr) => {
match $self {
crate::connector::pool_cache::SqlPool::Postgres($p) => {
let $rows_to_json = crate::connector::sql_decode::pg_rows_to_json;
let $bind = crate::connector::pool_cache::bind_params::<sqlx::Postgres>;
let $typed_args = crate::connector::sql_encode::pg_typed_args;
let $write_result = crate::connector::pool_cache::pg_write_result;
$body
}
crate::connector::pool_cache::SqlPool::MySql($p) => {
let $rows_to_json = crate::connector::sql_decode::mysql_rows_to_json;
let $bind = crate::connector::pool_cache::bind_params::<sqlx::MySql>;
let $typed_args = crate::connector::sql_encode::mysql_typed_args;
let $write_result = crate::connector::pool_cache::mysql_write_result;
$body
}
crate::connector::pool_cache::SqlPool::Sqlite($p) => {
let $rows_to_json = crate::connector::sql_decode::sqlite_rows_to_json;
let $bind = crate::connector::pool_cache::bind_params::<sqlx::Sqlite>;
let $typed_args = crate::connector::sql_encode::sqlite_typed_args;
let $write_result = crate::connector::pool_cache::sqlite_write_result;
$body
}
}
};
($self:expr, $p:ident, $rows_to_json:ident, $bind:ident, $write_result:ident => $body:expr) => {
crate::connector::pool_cache::dispatch_sql_pool!(
$self, $p, $rows_to_json, $bind, _typed_args, $write_result => $body
)
};
($self:expr, $p:ident, $rows_to_json:ident, $bind:ident => $body:expr) => {
crate::connector::pool_cache::dispatch_sql_pool!(
$self, $p, $rows_to_json, $bind, _typed_args, _write_result => $body
)
};
}
pub(crate) use dispatch_sql_pool;
impl SqlPool {
pub fn dialect(&self) -> crate::query::SqlDialect {
match self {
SqlPool::Postgres(_) => crate::query::SqlDialect::Postgres,
SqlPool::MySql(_) => crate::query::SqlDialect::Mysql,
SqlPool::Sqlite(_) => crate::query::SqlDialect::Sqlite,
}
}
pub fn reports_last_insert_id(&self) -> bool {
!matches!(self, SqlPool::Postgres(_))
}
pub async fn ping(&self) -> Result<(), sqlx::Error> {
dispatch_sql_pool!(self, p, _d, _b => sqlx::query("SELECT 1").execute(p).await.map(|_| ()))
}
pub fn is_closed(&self) -> bool {
dispatch_sql_pool!(self, p, _d, _b => p.is_closed())
}
async fn close(self) {
dispatch_sql_pool!(self, p, _d, _b => p.close().await)
}
}
pub struct SqlPoolCache {
cache: LruCache<SqlPool>,
}
impl SqlPoolCache {
pub fn new(max_entries: usize) -> Self {
Self {
cache: LruCache::with_evict_handler(max_entries, "sql_pool", |pool: SqlPool| {
tokio::spawn(async move { pool.close().await });
}),
}
}
pub async fn get_pool(
&self,
connector_name: &str,
config: &DbConnectorConfig,
) -> Result<SqlPool, OrionError> {
let conn_str = config.connection_string.clone();
let max_conns = config.max_connections.unwrap_or(5);
let connect_timeout = config.connect_timeout_ms.unwrap_or(5000);
self.cache
.get_or_create(connector_name, || async move {
crate::validation::check_db_endpoint(connector_name, config).await?;
let timeout = Duration::from_millis(connect_timeout);
let pool = match crate::storage::detect_backend(&conn_str)? {
crate::storage::DbBackend::Postgres => SqlPool::Postgres(
PgPoolOptions::new()
.max_connections(max_conns)
.acquire_timeout(timeout)
.connect(&conn_str)
.await
.map_err(|e| connect_failed(connector_name, e))?,
),
crate::storage::DbBackend::Mysql => SqlPool::MySql(
MySqlPoolOptions::new()
.max_connections(max_conns)
.acquire_timeout(timeout)
.connect(&conn_str)
.await
.map_err(|e| connect_failed(connector_name, e))?,
),
crate::storage::DbBackend::Sqlite => SqlPool::Sqlite(
SqlitePoolOptions::new()
.max_connections(max_conns)
.acquire_timeout(timeout)
.connect(&conn_str)
.await
.map_err(|e| connect_failed(connector_name, e))?,
),
};
Ok(pool)
})
.await
}
pub async fn evict(&self, connector_name: &str) {
self.cache.evict(connector_name).await;
}
pub async fn evict_all(&self) {
self.cache.evict_all().await;
}
}
impl Default for SqlPoolCache {
fn default() -> Self {
Self::new(100)
}
}
pub(crate) fn pg_write_result(r: &sqlx::postgres::PgQueryResult) -> (u64, Option<i64>) {
(r.rows_affected(), None)
}
pub(crate) fn mysql_write_result(r: &sqlx::mysql::MySqlQueryResult) -> (u64, Option<i64>) {
(r.rows_affected(), i64::try_from(r.last_insert_id()).ok())
}
pub(crate) fn sqlite_write_result(r: &sqlx::sqlite::SqliteQueryResult) -> (u64, Option<i64>) {
(r.rows_affected(), Some(r.last_insert_rowid()))
}
fn connect_failed(connector_name: &str, e: sqlx::Error) -> OrionError {
OrionError::Internal {
context: format!("Failed to connect to external DB '{connector_name}'"),
source: Some(Box::new(e)),
}
}
pub(crate) fn bind_params<'q, DB>(
mut query: sqlx::query::Query<'q, DB, <DB as sqlx::Database>::Arguments>,
params: &'q [serde_json::Value],
) -> sqlx::query::Query<'q, DB, <DB as sqlx::Database>::Arguments>
where
DB: sqlx::Database,
&'q str: sqlx::Encode<'q, DB> + sqlx::Type<DB>,
i64: sqlx::Encode<'q, DB> + sqlx::Type<DB>,
f64: sqlx::Encode<'q, DB> + sqlx::Type<DB>,
bool: sqlx::Encode<'q, DB> + sqlx::Type<DB>,
String: sqlx::Encode<'q, DB> + sqlx::Type<DB>,
Option<String>: sqlx::Encode<'q, DB> + sqlx::Type<DB>,
{
for param in params {
query = match param {
serde_json::Value::String(s) => query.bind(s.as_str()),
serde_json::Value::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(n.to_string())
}
}
serde_json::Value::Bool(b) => query.bind(*b),
serde_json::Value::Null => query.bind(None::<String>),
other => query.bind(other.to_string()),
};
}
query
}