sz_orm_sqlx/
unified_pool.rs1use 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
44pub struct UnifiedPool {
49 backend: AnyBackend,
50 pool: Pool,
51}
52
53impl UnifiedPool {
54 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 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 pub fn from_pool(pool: Pool, backend: AnyBackend) -> Self {
127 Self { backend, pool }
128 }
129
130 #[inline]
132 pub fn backend(&self) -> AnyBackend {
133 self.backend
134 }
135
136 #[inline]
138 pub fn dialect(&self) -> Box<dyn Dialect> {
139 self.backend.dialect()
140 }
141
142 #[inline]
144 pub async fn acquire(&self) -> Result<PooledConnection, PoolError> {
145 self.pool.acquire().await
146 }
147
148 #[inline]
150 pub fn resize(&self, new_max: usize) {
151 self.pool.resize(new_max);
152 }
153
154 #[inline]
156 pub async fn close_all(&self) {
157 self.pool.close_all().await;
158 }
159
160 #[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}