sz_orm_sqlx/
any_driver.rs1use std::future::Future;
32use std::pin::Pin;
33use std::sync::Arc;
34use sz_orm_core::{Connection, ConnectionFactory, DbError, QueryRows};
35
36use crate::any::{
37 MySqlPoolHandle, PgPoolHandle, SqlitePoolHandle, SqlxMySqlConnectionFactory,
38 SqlxPgConnectionFactory, SqlxSqliteConnectionFactory,
39};
40
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
43pub enum AnyBackend {
44 MySql,
46 Postgres,
48 Sqlite,
50}
51
52impl AnyBackend {
53 pub fn from_dsn(dsn: &str) -> Result<Self, DbError> {
61 if dsn.starts_with("mysql://") || dsn.starts_with("mariadb://") {
62 Ok(AnyBackend::MySql)
63 } else if dsn.starts_with("postgres://") || dsn.starts_with("postgresql://") {
64 Ok(AnyBackend::Postgres)
65 } else if dsn.starts_with("sqlite://") || dsn.starts_with("sqlite:") {
66 Ok(AnyBackend::Sqlite)
67 } else {
68 Err(DbError::ConnectionRefused(format!(
69 "未知的 DSN scheme: {}(支持 mysql/postgres/sqlite)",
70 dsn
71 )))
72 }
73 }
74
75 pub fn name(&self) -> &'static str {
77 match self {
78 AnyBackend::MySql => "mysql",
79 AnyBackend::Postgres => "postgres",
80 AnyBackend::Sqlite => "sqlite",
81 }
82 }
83}
84
85pub struct AnyPool {
87 backend: AnyBackend,
88 factory: Arc<dyn ConnectionFactory>,
89}
90
91impl AnyPool {
92 pub async fn connect(dsn: &str) -> Result<Self, DbError> {
99 let backend = AnyBackend::from_dsn(dsn)?;
100 let factory: Arc<dyn ConnectionFactory> = match backend {
101 AnyBackend::MySql => {
102 let handle = Arc::new(MySqlPoolHandle::connect(dsn).await?);
103 Arc::new(SqlxMySqlConnectionFactory::new(handle))
104 }
105 AnyBackend::Postgres => {
106 let handle = Arc::new(PgPoolHandle::connect(dsn).await?);
107 Arc::new(SqlxPgConnectionFactory::new(handle))
108 }
109 AnyBackend::Sqlite => {
110 let handle = Arc::new(SqlitePoolHandle::connect(dsn).await?);
111 Arc::new(SqlxSqliteConnectionFactory::new(handle))
112 }
113 };
114 Ok(Self { backend, factory })
115 }
116
117 pub fn from_factory(backend: AnyBackend, factory: Arc<dyn ConnectionFactory>) -> Self {
119 Self { backend, factory }
120 }
121
122 pub fn backend(&self) -> AnyBackend {
124 self.backend
125 }
126
127 pub async fn create(&self) -> Result<AnyConnection, DbError> {
129 let conn = self.factory.create().await?;
130 Ok(AnyConnection {
131 backend: self.backend,
132 inner: conn,
133 })
134 }
135}
136
137pub struct AnyConnection {
139 backend: AnyBackend,
140 inner: Box<dyn Connection>,
141}
142
143impl AnyConnection {
144 pub fn backend(&self) -> AnyBackend {
146 self.backend
147 }
148}
149
150impl Connection for AnyConnection {
151 fn execute<'a>(
152 &'a mut self,
153 sql: &'a str,
154 ) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
155 self.inner.execute(sql)
156 }
157
158 fn query<'a>(
159 &'a mut self,
160 sql: &'a str,
161 ) -> Pin<Box<dyn Future<Output = Result<QueryRows, DbError>> + Send + 'a>> {
162 self.inner.query(sql)
163 }
164
165 fn begin_transaction<'a>(
166 &'a mut self,
167 ) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
168 self.inner.begin_transaction()
169 }
170
171 fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
172 self.inner.commit()
173 }
174
175 fn rollback<'a>(
176 &'a mut self,
177 ) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
178 self.inner.rollback()
179 }
180
181 fn is_connected(&self) -> bool {
182 self.inner.is_connected()
183 }
184
185 fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
186 self.inner.ping()
187 }
188
189 fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
190 self.inner.close()
191 }
192}
193
194#[cfg(test)]
195mod tests {
196 use super::*;
197
198 #[test]
199 fn test_any_backend_from_dsn_mysql() {
200 assert_eq!(
201 AnyBackend::from_dsn("mysql://root:pass@127.0.0.1/db").unwrap(),
202 AnyBackend::MySql
203 );
204 assert_eq!(
205 AnyBackend::from_dsn("mariadb://root:pass@127.0.0.1/db").unwrap(),
206 AnyBackend::MySql
207 );
208 }
209
210 #[test]
211 fn test_any_backend_from_dsn_postgres() {
212 assert_eq!(
213 AnyBackend::from_dsn("postgres://user:pass@127.0.0.1/db").unwrap(),
214 AnyBackend::Postgres
215 );
216 assert_eq!(
217 AnyBackend::from_dsn("postgresql://user:pass@127.0.0.1/db").unwrap(),
218 AnyBackend::Postgres
219 );
220 }
221
222 #[test]
223 fn test_any_backend_from_dsn_sqlite() {
224 assert_eq!(
225 AnyBackend::from_dsn("sqlite::memory:").unwrap(),
226 AnyBackend::Sqlite
227 );
228 assert_eq!(
229 AnyBackend::from_dsn("sqlite://./test.db").unwrap(),
230 AnyBackend::Sqlite
231 );
232 }
233
234 #[test]
235 fn test_any_backend_from_dsn_unknown() {
236 let result = AnyBackend::from_dsn("redis://127.0.0.1");
237 assert!(result.is_err());
238 if let Err(e) = result {
239 let msg = format!("{}", e);
240 assert!(msg.contains("redis") || msg.contains("未知"));
241 }
242 }
243
244 #[test]
245 fn test_any_backend_name() {
246 assert_eq!(AnyBackend::MySql.name(), "mysql");
247 assert_eq!(AnyBackend::Postgres.name(), "postgres");
248 assert_eq!(AnyBackend::Sqlite.name(), "sqlite");
249 }
250
251 #[test]
252 fn test_any_backend_equality() {
253 assert_eq!(AnyBackend::MySql, AnyBackend::MySql);
254 assert_ne!(AnyBackend::MySql, AnyBackend::Postgres);
255 assert_ne!(AnyBackend::Postgres, AnyBackend::Sqlite);
256 }
257
258 #[tokio::test]
261 async fn test_any_pool_sqlite_connect_and_query() {
262 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
263 assert_eq!(pool.backend(), AnyBackend::Sqlite);
264
265 let mut conn = pool.create().await.unwrap();
266 assert_eq!(conn.backend(), AnyBackend::Sqlite);
267
268 conn.execute("CREATE TABLE test_any (id INTEGER PRIMARY KEY, name TEXT)")
270 .await
271 .unwrap();
272 conn.execute("INSERT INTO test_any (name) VALUES ('Alice')")
273 .await
274 .unwrap();
275 conn.execute("INSERT INTO test_any (name) VALUES ('Bob')")
276 .await
277 .unwrap();
278
279 let rows = conn
281 .query("SELECT * FROM test_any ORDER BY id")
282 .await
283 .unwrap();
284 assert_eq!(rows.len(), 2);
285 assert_eq!(rows[0].get("name").and_then(|v| v.as_str()), Some("Alice"));
286 assert_eq!(rows[1].get("name").and_then(|v| v.as_str()), Some("Bob"));
287 }
288
289 #[tokio::test]
290 async fn test_any_pool_sqlite_transaction_commit() {
291 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
292 let mut conn = pool.create().await.unwrap();
293
294 conn.execute("CREATE TABLE tx_test (id INTEGER PRIMARY KEY, val INTEGER)")
295 .await
296 .unwrap();
297
298 conn.begin_transaction().await.unwrap();
300 conn.execute("INSERT INTO tx_test (val) VALUES (1)")
301 .await
302 .unwrap();
303 conn.execute("INSERT INTO tx_test (val) VALUES (2)")
304 .await
305 .unwrap();
306 conn.commit().await.unwrap();
307
308 let rows = conn.query("SELECT * FROM tx_test").await.unwrap();
309 assert_eq!(rows.len(), 2);
310 }
311
312 #[tokio::test]
313 async fn test_any_pool_sqlite_transaction_rollback() {
314 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
315 let mut conn = pool.create().await.unwrap();
316
317 conn.execute("CREATE TABLE tx_rb (id INTEGER PRIMARY KEY, val INTEGER)")
318 .await
319 .unwrap();
320
321 conn.begin_transaction().await.unwrap();
322 conn.execute("INSERT INTO tx_rb (val) VALUES (1)")
323 .await
324 .unwrap();
325 conn.rollback().await.unwrap();
326
327 let rows = conn.query("SELECT * FROM tx_rb").await.unwrap();
328 assert_eq!(rows.len(), 0);
329 }
330
331 #[tokio::test]
332 async fn test_any_pool_sqlite_ping() {
333 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
334 let mut conn = pool.create().await.unwrap();
335 let ok = conn.ping().await;
336 assert!(ok);
337 assert!(conn.is_connected());
338 }
339
340 #[tokio::test]
341 async fn test_any_pool_invalid_dsn() {
342 let result = AnyPool::connect("invalid://dsn").await;
343 assert!(result.is_err());
344 }
345
346 #[tokio::test]
347 async fn test_any_pool_sqlite_count_query() {
348 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
349 let mut conn = pool.create().await.unwrap();
350
351 conn.execute("CREATE TABLE cnt (id INTEGER PRIMARY KEY)")
352 .await
353 .unwrap();
354 for i in 1..=5 {
355 conn.execute(&format!("INSERT INTO cnt (id) VALUES ({})", i))
356 .await
357 .unwrap();
358 }
359
360 let rows = conn.query("SELECT * FROM cnt").await.unwrap();
362 assert_eq!(rows.len(), 5);
363 }
364}