Skip to main content

sz_orm_sqlx/
unified_pool.rs

1//! UnifiedPool — 统一连接池抽象(v2.2.0 A-3)
2//!
3//! 包装 `sz_orm_core::Pool`(完整连接池)+ `AnyBackend`,提供 5 后端透明的统一类型。
4//! 供 sz-rust AppState 持有单一类型 `Arc<UnifiedPool>`,业务代码无需感知后端类型。
5//!
6//! # 设计
7//!
8//! - `UnifiedPool` 是 `Pool` 的 newtype 包装,所有方法委托 `Pool`(零能力丢失)
9//! - `from_pool` 提供零成本迁移路径:sz-rust 从 `Arc<Pool>` 迁移到 `Arc<UnifiedPool>`
10//! - `connect`/`connect_with_config` 根据 DSN 自动识别后端并创建完整连接池
11//!
12//! # 用法
13//!
14//! ```ignore
15//! use sz_orm_sqlx::UnifiedPool;
16//!
17//! // 从 DSN 自动识别后端,创建完整连接池
18//! let pool = UnifiedPool::connect("mysql://root:pass@127.0.0.1/db").await?;
19//! let mut conn = pool.acquire().await?;
20//!
21//! // 运行时切换后端
22//! let pg_pool = UnifiedPool::connect("postgres://user:pass@127.0.0.1/db").await?;
23//! let d = pg_pool.dialect(); // 自动返回 PostgreSqlDialect
24//! ```
25
26use std::sync::Arc;
27use sz_orm_core::{
28    ConnectionFactory, DbError, Dialect, Pool, PoolConfig, PoolConfigBuilder, PoolError,
29    PoolStatus, PooledConnection,
30};
31
32use crate::any::{
33    MySqlPoolHandle, PgPoolHandle, SqlitePoolHandle, SqlxMySqlConnectionFactory,
34    SqlxPgConnectionFactory, SqlxSqliteConnectionFactory,
35};
36use crate::any_driver::AnyBackend;
37
38#[cfg(feature = "oracle")]
39use sz_orm_oracle::{OracleConnectionFactory, OraclePoolHandle};
40
41#[cfg(feature = "mssql")]
42use sz_orm_mssql::{MssqlConnectionFactory, MssqlPoolHandle};
43
44/// 统一连接池:包装 `Pool` + `AnyBackend`,5 后端透明切换(v2.2.0 新增)
45///
46/// 供 sz-rust AppState 持有 `Arc<UnifiedPool>`,业务代码无需感知后端类型。
47/// 所有方法委托内部 `Pool`,零能力丢失。
48pub struct UnifiedPool {
49    backend: AnyBackend,
50    pool: Pool,
51}
52
53impl UnifiedPool {
54    /// 连接数据库,根据 DSN 自动识别后端,使用默认 PoolConfig
55    ///
56    /// 默认配置:max_size=10, timeout=30s
57    pub async fn connect(dsn: &str) -> Result<Self, DbError> {
58        let config = PoolConfigBuilder::new()
59            .build()
60            .map_err(DbError::PoolError)?;
61        Self::connect_with_config(dsn, config).await
62    }
63
64    /// 连接数据库,根据 DSN 自动识别后端,使用自定义 PoolConfig
65    pub async fn connect_with_config(dsn: &str, config: PoolConfig) -> Result<Self, DbError> {
66        let backend = AnyBackend::from_dsn(dsn)?;
67        let factory: Arc<dyn ConnectionFactory> = match backend {
68            AnyBackend::MySql => {
69                let handle = Arc::new(MySqlPoolHandle::connect(dsn).await?);
70                Arc::new(SqlxMySqlConnectionFactory::new(handle))
71            }
72            AnyBackend::Postgres => {
73                let handle = Arc::new(PgPoolHandle::connect(dsn).await?);
74                Arc::new(SqlxPgConnectionFactory::new(handle))
75            }
76            AnyBackend::Sqlite => {
77                let handle = Arc::new(SqlitePoolHandle::connect(dsn).await?);
78                Arc::new(SqlxSqliteConnectionFactory::new(handle))
79            }
80            AnyBackend::Oracle => {
81                #[cfg(feature = "oracle")]
82                {
83                    let (username, password, connect_string) =
84                        crate::any_driver::parse_oracle_dsn(dsn)?;
85                    let handle = Arc::new(OraclePoolHandle::connect(
86                        &username,
87                        &password,
88                        &connect_string,
89                    )?);
90                    Arc::new(OracleConnectionFactory::new(handle))
91                }
92                #[cfg(not(feature = "oracle"))]
93                {
94                    return Err(DbError::ConnectionRefused(
95                        "Oracle 后端未启用,请在 Cargo.toml 中添加 features = [\"oracle\"]"
96                            .to_string(),
97                    ));
98                }
99            }
100            AnyBackend::Mssql => {
101                #[cfg(feature = "mssql")]
102                {
103                    let ado_string = crate::any_driver::parse_mssql_dsn(dsn)?;
104                    let handle = Arc::new(MssqlPoolHandle::connect(&ado_string).await?);
105                    Arc::new(MssqlConnectionFactory::new(handle))
106                }
107                #[cfg(not(feature = "mssql"))]
108                {
109                    return Err(DbError::ConnectionRefused(
110                        "MSSQL 后端未启用,请在 Cargo.toml 中添加 features = [\"mssql\"]"
111                            .to_string(),
112                    ));
113                }
114            }
115        };
116        let pool = Pool::new(config, factory).map_err(DbError::PoolError)?;
117        Ok(Self { backend, pool })
118    }
119
120    /// 从已有的 Pool 构造 UnifiedPool(零成本迁移)
121    ///
122    /// 供 sz-rust 从 `Arc<Pool>` 迁移到 `Arc<UnifiedPool>`:
123    /// ```ignore
124    /// let unified = UnifiedPool::from_pool(existing_pool, AnyBackend::MySql);
125    /// ```
126    pub fn from_pool(pool: Pool, backend: AnyBackend) -> Self {
127        Self { backend, pool }
128    }
129
130    /// 获取后端类型
131    #[inline]
132    pub fn backend(&self) -> AnyBackend {
133        self.backend
134    }
135
136    /// 返回对应后端的 Dialect 实例
137    #[inline]
138    pub fn dialect(&self) -> Box<dyn Dialect> {
139        self.backend.dialect()
140    }
141
142    /// 获取连接(委托 Pool::acquire)
143    #[inline]
144    pub async fn acquire(&self) -> Result<PooledConnection, PoolError> {
145        self.pool.acquire().await
146    }
147
148    /// 调整连接池大小(委托 Pool::resize)
149    #[inline]
150    pub fn resize(&self, new_max: usize) {
151        self.pool.resize(new_max);
152    }
153
154    /// 关闭所有连接(委托 Pool::close_all)
155    #[inline]
156    pub async fn close_all(&self) {
157        self.pool.close_all().await;
158    }
159
160    /// 获取连接池状态(委托 Pool::status)
161    #[inline]
162    pub async fn status(&self) -> PoolStatus {
163        self.pool.status().await
164    }
165}
166
167impl std::fmt::Debug for UnifiedPool {
168    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
169        f.debug_struct("UnifiedPool")
170            .field("backend", &self.backend)
171            .finish_non_exhaustive()
172    }
173}
174
175#[cfg(test)]
176mod tests {
177    use super::*;
178
179    #[tokio::test]
180    async fn test_unified_pool_sqlite_connect() {
181        let pool = UnifiedPool::connect("sqlite::memory:").await.unwrap();
182        assert_eq!(pool.backend(), AnyBackend::Sqlite);
183
184        let mut conn = pool.acquire().await.unwrap();
185        conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY)")
186            .await
187            .unwrap();
188        conn.execute("INSERT INTO t (id) VALUES (1)").await.unwrap();
189        let rows = conn.query("SELECT * FROM t").await.unwrap();
190        assert_eq!(rows.len(), 1);
191    }
192
193    #[tokio::test]
194    async fn test_unified_pool_dialect() {
195        let pool = UnifiedPool::connect("sqlite::memory:").await.unwrap();
196        let d = pool.dialect();
197        assert_eq!(d.db_type(), sz_orm_core::DbType::Sqlite);
198    }
199
200    #[tokio::test]
201    async fn test_unified_pool_from_pool() {
202        let handle = Arc::new(SqlitePoolHandle::connect("sqlite::memory:").await.unwrap());
203        let factory = Arc::new(SqlxSqliteConnectionFactory::new(handle));
204        let config = PoolConfigBuilder::new().build().unwrap();
205        let pool = Pool::new(config, factory).unwrap();
206
207        let unified = UnifiedPool::from_pool(pool, AnyBackend::Sqlite);
208        assert_eq!(unified.backend(), AnyBackend::Sqlite);
209
210        let mut conn = unified.acquire().await.unwrap();
211        conn.execute("SELECT 1").await.unwrap();
212    }
213
214    #[tokio::test]
215    async fn test_unified_pool_resize_and_close() {
216        let pool = UnifiedPool::connect("sqlite::memory:").await.unwrap();
217        pool.resize(20);
218        let status = pool.status().await;
219        assert_eq!(status.max, 20);
220        pool.close_all().await;
221    }
222
223    #[tokio::test]
224    async fn test_unified_pool_invalid_dsn() {
225        let result = UnifiedPool::connect("invalid://dsn").await;
226        assert!(result.is_err());
227    }
228}