use std::sync::Arc;
use sz_orm_core::{
ConnectionFactory, DbError, Dialect, Pool, PoolConfig, PoolConfigBuilder, PoolError,
PoolStatus, PooledConnection,
};
use crate::any::{
MySqlPoolHandle, PgPoolHandle, SqlitePoolHandle, SqlxMySqlConnectionFactory,
SqlxPgConnectionFactory, SqlxSqliteConnectionFactory,
};
use crate::any_driver::AnyBackend;
#[cfg(feature = "oracle")]
use sz_orm_oracle::{OracleConnectionFactory, OraclePoolHandle};
#[cfg(feature = "mssql")]
use sz_orm_mssql::{MssqlConnectionFactory, MssqlPoolHandle};
pub struct UnifiedPool {
backend: AnyBackend,
pool: Pool,
}
impl UnifiedPool {
pub async fn connect(dsn: &str) -> Result<Self, DbError> {
let config = PoolConfigBuilder::new()
.build()
.map_err(DbError::PoolError)?;
Self::connect_with_config(dsn, config).await
}
pub async fn connect_with_config(dsn: &str, config: PoolConfig) -> Result<Self, DbError> {
let backend = AnyBackend::from_dsn(dsn)?;
let factory: Arc<dyn ConnectionFactory> = match backend {
AnyBackend::MySql => {
let handle = Arc::new(MySqlPoolHandle::connect(dsn).await?);
Arc::new(SqlxMySqlConnectionFactory::new(handle))
}
AnyBackend::Postgres => {
let handle = Arc::new(PgPoolHandle::connect(dsn).await?);
Arc::new(SqlxPgConnectionFactory::new(handle))
}
AnyBackend::Sqlite => {
let handle = Arc::new(SqlitePoolHandle::connect(dsn).await?);
Arc::new(SqlxSqliteConnectionFactory::new(handle))
}
AnyBackend::Oracle => {
#[cfg(feature = "oracle")]
{
let (username, password, connect_string) =
crate::any_driver::parse_oracle_dsn(dsn)?;
let handle = Arc::new(OraclePoolHandle::connect(
&username,
&password,
&connect_string,
)?);
Arc::new(OracleConnectionFactory::new(handle))
}
#[cfg(not(feature = "oracle"))]
{
return Err(DbError::ConnectionRefused(
"Oracle 后端未启用,请在 Cargo.toml 中添加 features = [\"oracle\"]"
.to_string(),
));
}
}
AnyBackend::Mssql => {
#[cfg(feature = "mssql")]
{
let ado_string = crate::any_driver::parse_mssql_dsn(dsn)?;
let handle = Arc::new(MssqlPoolHandle::connect(&ado_string).await?);
Arc::new(MssqlConnectionFactory::new(handle))
}
#[cfg(not(feature = "mssql"))]
{
return Err(DbError::ConnectionRefused(
"MSSQL 后端未启用,请在 Cargo.toml 中添加 features = [\"mssql\"]"
.to_string(),
));
}
}
};
let pool = Pool::new(config, factory).map_err(DbError::PoolError)?;
Ok(Self { backend, pool })
}
pub fn from_pool(pool: Pool, backend: AnyBackend) -> Self {
Self { backend, pool }
}
#[inline]
pub fn backend(&self) -> AnyBackend {
self.backend
}
#[inline]
pub fn dialect(&self) -> Box<dyn Dialect> {
self.backend.dialect()
}
#[inline]
pub async fn acquire(&self) -> Result<PooledConnection, PoolError> {
self.pool.acquire().await
}
#[inline]
pub fn resize(&self, new_max: usize) {
self.pool.resize(new_max);
}
#[inline]
pub async fn close_all(&self) {
self.pool.close_all().await;
}
#[inline]
pub async fn status(&self) -> PoolStatus {
self.pool.status().await
}
}
impl std::fmt::Debug for UnifiedPool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UnifiedPool")
.field("backend", &self.backend)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_unified_pool_sqlite_connect() {
let pool = UnifiedPool::connect("sqlite::memory:").await.unwrap();
assert_eq!(pool.backend(), AnyBackend::Sqlite);
let mut conn = pool.acquire().await.unwrap();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY)")
.await
.unwrap();
conn.execute("INSERT INTO t (id) VALUES (1)").await.unwrap();
let rows = conn.query("SELECT * FROM t").await.unwrap();
assert_eq!(rows.len(), 1);
}
#[tokio::test]
async fn test_unified_pool_dialect() {
let pool = UnifiedPool::connect("sqlite::memory:").await.unwrap();
let d = pool.dialect();
assert_eq!(d.db_type(), sz_orm_core::DbType::Sqlite);
}
#[tokio::test]
async fn test_unified_pool_from_pool() {
let handle = Arc::new(SqlitePoolHandle::connect("sqlite::memory:").await.unwrap());
let factory = Arc::new(SqlxSqliteConnectionFactory::new(handle));
let config = PoolConfigBuilder::new().build().unwrap();
let pool = Pool::new(config, factory).unwrap();
let unified = UnifiedPool::from_pool(pool, AnyBackend::Sqlite);
assert_eq!(unified.backend(), AnyBackend::Sqlite);
let mut conn = unified.acquire().await.unwrap();
conn.execute("SELECT 1").await.unwrap();
}
#[tokio::test]
async fn test_unified_pool_resize_and_close() {
let pool = UnifiedPool::connect("sqlite::memory:").await.unwrap();
pool.resize(20);
let status = pool.status().await;
assert_eq!(status.max, 20);
pool.close_all().await;
}
#[tokio::test]
async fn test_unified_pool_invalid_dsn() {
let result = UnifiedPool::connect("invalid://dsn").await;
assert!(result.is_err());
}
}