1use std::future::Future;
32use std::pin::Pin;
33use std::sync::Arc;
34use sz_orm_core::{
35 Connection, ConnectionFactory, DbError, Dialect, MySqlDialect, OracleDialect,
36 PostgreSqlDialect, QueryRows, SqlServerDialect, SqliteDialect,
37};
38
39use crate::any::{
40 MySqlPoolHandle, PgPoolHandle, SqlitePoolHandle, SqlxMySqlConnectionFactory,
41 SqlxPgConnectionFactory, SqlxSqliteConnectionFactory,
42};
43
44#[cfg(feature = "oracle")]
45use sz_orm_oracle::{OracleConnectionFactory, OraclePoolHandle};
46
47#[cfg(feature = "mssql")]
48use sz_orm_mssql::{MssqlConnectionFactory, MssqlPoolHandle};
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
56#[non_exhaustive]
57pub enum AnyBackend {
58 MySql,
60 Postgres,
62 Sqlite,
64 Oracle,
66 Mssql,
68}
69
70impl AnyBackend {
71 pub fn from_dsn(dsn: &str) -> Result<Self, DbError> {
81 if dsn.starts_with("mysql://") || dsn.starts_with("mariadb://") {
82 Ok(AnyBackend::MySql)
83 } else if dsn.starts_with("postgres://") || dsn.starts_with("postgresql://") {
84 Ok(AnyBackend::Postgres)
85 } else if dsn.starts_with("sqlite://") || dsn.starts_with("sqlite:") {
86 Ok(AnyBackend::Sqlite)
87 } else if dsn.starts_with("oracle://") {
88 Ok(AnyBackend::Oracle)
89 } else if dsn.starts_with("mssql://") || dsn.starts_with("sqlserver://") {
90 Ok(AnyBackend::Mssql)
91 } else {
92 Err(DbError::ConnectionRefused(format!(
93 "未知的 DSN scheme: {}(支持 mysql/postgres/sqlite/oracle/mssql)",
94 dsn
95 )))
96 }
97 }
98
99 pub fn name(&self) -> &'static str {
101 match self {
102 AnyBackend::MySql => "mysql",
103 AnyBackend::Postgres => "postgres",
104 AnyBackend::Sqlite => "sqlite",
105 AnyBackend::Oracle => "oracle",
106 AnyBackend::Mssql => "mssql",
107 }
108 }
109
110 pub fn dialect(&self) -> Box<dyn Dialect> {
118 match self {
119 AnyBackend::MySql => Box::new(MySqlDialect),
120 AnyBackend::Postgres => Box::new(PostgreSqlDialect),
121 AnyBackend::Sqlite => Box::new(SqliteDialect),
122 AnyBackend::Oracle => Box::new(OracleDialect),
123 AnyBackend::Mssql => Box::new(SqlServerDialect),
124 }
125 }
126}
127
128pub struct AnyPool {
130 backend: AnyBackend,
131 factory: Arc<dyn ConnectionFactory>,
132}
133
134impl AnyPool {
135 pub async fn connect(dsn: &str) -> Result<Self, DbError> {
143 let backend = AnyBackend::from_dsn(dsn)?;
144 let factory: Arc<dyn ConnectionFactory> = match backend {
145 AnyBackend::MySql => {
146 let handle = Arc::new(MySqlPoolHandle::connect(dsn).await?);
147 Arc::new(SqlxMySqlConnectionFactory::new(handle))
148 }
149 AnyBackend::Postgres => {
150 let handle = Arc::new(PgPoolHandle::connect(dsn).await?);
151 Arc::new(SqlxPgConnectionFactory::new(handle))
152 }
153 AnyBackend::Sqlite => {
154 let handle = Arc::new(SqlitePoolHandle::connect(dsn).await?);
155 Arc::new(SqlxSqliteConnectionFactory::new(handle))
156 }
157 AnyBackend::Oracle => {
158 #[cfg(feature = "oracle")]
159 {
160 let (username, password, connect_string) = parse_oracle_dsn(dsn)?;
161 let handle = Arc::new(OraclePoolHandle::connect(
162 &username,
163 &password,
164 &connect_string,
165 )?);
166 Arc::new(OracleConnectionFactory::new(handle))
167 }
168 #[cfg(not(feature = "oracle"))]
169 {
170 return Err(DbError::ConnectionRefused(
171 "Oracle 后端未启用,请在 Cargo.toml 中添加 features = [\"oracle\"]"
172 .to_string(),
173 ));
174 }
175 }
176 AnyBackend::Mssql => {
177 #[cfg(feature = "mssql")]
178 {
179 let ado_string = parse_mssql_dsn(dsn)?;
180 let handle = Arc::new(MssqlPoolHandle::connect(&ado_string).await?);
181 Arc::new(MssqlConnectionFactory::new(handle))
182 }
183 #[cfg(not(feature = "mssql"))]
184 {
185 return Err(DbError::ConnectionRefused(
186 "MSSQL 后端未启用,请在 Cargo.toml 中添加 features = [\"mssql\"]"
187 .to_string(),
188 ));
189 }
190 }
191 };
192 Ok(Self { backend, factory })
193 }
194
195 pub fn from_factory(backend: AnyBackend, factory: Arc<dyn ConnectionFactory>) -> Self {
197 Self { backend, factory }
198 }
199
200 pub fn backend(&self) -> AnyBackend {
202 self.backend
203 }
204
205 pub fn dialect(&self) -> Box<dyn Dialect> {
209 self.backend.dialect()
210 }
211
212 pub async fn create(&self) -> Result<AnyConnection, DbError> {
214 let conn = self.factory.create().await?;
215 Ok(AnyConnection {
216 backend: self.backend,
217 inner: conn,
218 })
219 }
220}
221
222#[allow(dead_code)]
228pub(crate) fn parse_oracle_dsn(dsn: &str) -> Result<(String, String, String), DbError> {
229 let rest = dsn
230 .strip_prefix("oracle://")
231 .ok_or_else(|| DbError::ConnectionRefused(format!("无效的 Oracle DSN: {}", dsn)))?;
232 parse_user_pass_host(rest, "Oracle")
233}
234
235#[allow(dead_code)]
241pub(crate) fn parse_mssql_dsn(dsn: &str) -> Result<String, DbError> {
242 let rest = dsn
243 .strip_prefix("mssql://")
244 .or_else(|| dsn.strip_prefix("sqlserver://"))
245 .ok_or_else(|| DbError::ConnectionRefused(format!("无效的 MSSQL DSN: {}", dsn)))?;
246 let (username, password, host_port_db) = parse_user_pass_host(rest, "MSSQL")?;
247 let (host_port, database) = host_port_db
248 .split_once('/')
249 .ok_or_else(|| DbError::ConnectionRefused(format!("MSSQL DSN 缺少 database: {}", dsn)))?;
250 let (host, port) = host_port.split_once(':').unwrap_or((host_port, "1433"));
251 Ok(format!(
252 "Server={},{};Database={};User Id={};Password={};",
253 host, port, database, username, password
254 ))
255}
256
257#[allow(dead_code)]
259fn parse_user_pass_host(rest: &str, backend: &str) -> Result<(String, String, String), DbError> {
260 let (userinfo, hostinfo) = rest
261 .split_once('@')
262 .ok_or_else(|| DbError::ConnectionRefused(format!("{} DSN 缺少 @: {}", backend, rest)))?;
263 let (username, password) = userinfo.split_once(':').ok_or_else(|| {
264 DbError::ConnectionRefused(format!("{} DSN 缺少 password: {}", backend, rest))
265 })?;
266 Ok((
267 username.to_string(),
268 password.to_string(),
269 hostinfo.to_string(),
270 ))
271}
272
273pub struct AnyConnection {
275 backend: AnyBackend,
276 inner: Box<dyn Connection>,
277}
278
279impl AnyConnection {
280 pub fn backend(&self) -> AnyBackend {
282 self.backend
283 }
284}
285
286impl Connection for AnyConnection {
287 fn execute<'a>(
288 &'a mut self,
289 sql: &'a str,
290 ) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
291 self.inner.execute(sql)
292 }
293
294 fn query<'a>(
295 &'a mut self,
296 sql: &'a str,
297 ) -> Pin<Box<dyn Future<Output = Result<QueryRows, DbError>> + Send + 'a>> {
298 self.inner.query(sql)
299 }
300
301 fn begin_transaction<'a>(
302 &'a mut self,
303 ) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
304 self.inner.begin_transaction()
305 }
306
307 fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
308 self.inner.commit()
309 }
310
311 fn rollback<'a>(
312 &'a mut self,
313 ) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
314 self.inner.rollback()
315 }
316
317 fn is_connected(&self) -> bool {
318 self.inner.is_connected()
319 }
320
321 fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
322 self.inner.ping()
323 }
324
325 fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
326 self.inner.close()
327 }
328}
329
330#[cfg(test)]
331mod tests {
332 use super::*;
333
334 #[test]
335 fn test_any_backend_from_dsn_mysql() {
336 assert_eq!(
337 AnyBackend::from_dsn("mysql://root:pass@127.0.0.1/db").unwrap(),
338 AnyBackend::MySql
339 );
340 assert_eq!(
341 AnyBackend::from_dsn("mariadb://root:pass@127.0.0.1/db").unwrap(),
342 AnyBackend::MySql
343 );
344 }
345
346 #[test]
347 fn test_any_backend_from_dsn_postgres() {
348 assert_eq!(
349 AnyBackend::from_dsn("postgres://user:pass@127.0.0.1/db").unwrap(),
350 AnyBackend::Postgres
351 );
352 assert_eq!(
353 AnyBackend::from_dsn("postgresql://user:pass@127.0.0.1/db").unwrap(),
354 AnyBackend::Postgres
355 );
356 }
357
358 #[test]
359 fn test_any_backend_from_dsn_sqlite() {
360 assert_eq!(
361 AnyBackend::from_dsn("sqlite::memory:").unwrap(),
362 AnyBackend::Sqlite
363 );
364 assert_eq!(
365 AnyBackend::from_dsn("sqlite://./test.db").unwrap(),
366 AnyBackend::Sqlite
367 );
368 }
369
370 #[test]
371 fn test_any_backend_from_dsn_oracle() {
372 assert_eq!(
373 AnyBackend::from_dsn("oracle://sys:test123@127.0.0.1:1521/freepdb1").unwrap(),
374 AnyBackend::Oracle
375 );
376 }
377
378 #[test]
379 fn test_any_backend_from_dsn_mssql() {
380 assert_eq!(
381 AnyBackend::from_dsn("mssql://sa:test123@127.0.0.1:1433/testdb").unwrap(),
382 AnyBackend::Mssql
383 );
384 }
385
386 #[test]
387 fn test_any_backend_from_dsn_sqlserver() {
388 assert_eq!(
389 AnyBackend::from_dsn("sqlserver://sa:test123@127.0.0.1:1433/testdb").unwrap(),
390 AnyBackend::Mssql
391 );
392 }
393
394 #[test]
395 fn test_any_backend_from_dsn_unknown_v22() {
396 let result = AnyBackend::from_dsn("redis://127.0.0.1");
397 assert!(result.is_err());
398 if let Err(e) = result {
399 let msg = format!("{}", e);
400 assert!(msg.contains("mysql"));
401 assert!(msg.contains("postgres"));
402 assert!(msg.contains("sqlite"));
403 assert!(msg.contains("oracle"));
404 assert!(msg.contains("mssql"));
405 }
406 }
407
408 #[test]
409 fn test_any_backend_name_v22() {
410 assert_eq!(AnyBackend::MySql.name(), "mysql");
411 assert_eq!(AnyBackend::Postgres.name(), "postgres");
412 assert_eq!(AnyBackend::Sqlite.name(), "sqlite");
413 assert_eq!(AnyBackend::Oracle.name(), "oracle");
414 assert_eq!(AnyBackend::Mssql.name(), "mssql");
415 }
416
417 #[test]
418 fn test_any_backend_equality() {
419 assert_eq!(AnyBackend::MySql, AnyBackend::MySql);
420 assert_ne!(AnyBackend::MySql, AnyBackend::Postgres);
421 assert_ne!(AnyBackend::Postgres, AnyBackend::Sqlite);
422 assert_ne!(AnyBackend::Oracle, AnyBackend::Mssql);
423 assert_ne!(AnyBackend::Oracle, AnyBackend::MySql);
424 }
425
426 #[test]
427 fn test_parse_oracle_dsn() {
428 let (user, pass, cs) =
429 parse_oracle_dsn("oracle://sys:test123@127.0.0.1:1521/freepdb1").unwrap();
430 assert_eq!(user, "sys");
431 assert_eq!(pass, "test123");
432 assert_eq!(cs, "127.0.0.1:1521/freepdb1");
433 }
434
435 #[test]
436 fn test_parse_mssql_dsn() {
437 let ado = parse_mssql_dsn("mssql://sa:test123@127.0.0.1:1433/testdb").unwrap();
438 assert!(ado.contains("Server=127.0.0.1,1433"));
439 assert!(ado.contains("Database=testdb"));
440 assert!(ado.contains("User Id=sa"));
441 assert!(ado.contains("Password=test123"));
442 }
443
444 #[test]
445 fn test_parse_mssql_dsn_sqlserver_scheme() {
446 let ado = parse_mssql_dsn("sqlserver://sa:pass@localhost/testdb").unwrap();
447 assert!(ado.contains("Server=localhost,1433"));
448 assert!(ado.contains("Database=testdb"));
449 }
450
451 #[tokio::test]
454 async fn test_any_pool_sqlite_connect_and_query() {
455 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
456 assert_eq!(pool.backend(), AnyBackend::Sqlite);
457
458 let mut conn = pool.create().await.unwrap();
459 assert_eq!(conn.backend(), AnyBackend::Sqlite);
460
461 conn.execute("CREATE TABLE test_any (id INTEGER PRIMARY KEY, name TEXT)")
463 .await
464 .unwrap();
465 conn.execute("INSERT INTO test_any (name) VALUES ('Alice')")
466 .await
467 .unwrap();
468 conn.execute("INSERT INTO test_any (name) VALUES ('Bob')")
469 .await
470 .unwrap();
471
472 let rows = conn
474 .query("SELECT * FROM test_any ORDER BY id")
475 .await
476 .unwrap();
477 assert_eq!(rows.len(), 2);
478 assert_eq!(rows[0].get("name").and_then(|v| v.as_str()), Some("Alice"));
479 assert_eq!(rows[1].get("name").and_then(|v| v.as_str()), Some("Bob"));
480 }
481
482 #[tokio::test]
483 async fn test_any_pool_sqlite_transaction_commit() {
484 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
485 let mut conn = pool.create().await.unwrap();
486
487 conn.execute("CREATE TABLE tx_test (id INTEGER PRIMARY KEY, val INTEGER)")
488 .await
489 .unwrap();
490
491 conn.begin_transaction().await.unwrap();
493 conn.execute("INSERT INTO tx_test (val) VALUES (1)")
494 .await
495 .unwrap();
496 conn.execute("INSERT INTO tx_test (val) VALUES (2)")
497 .await
498 .unwrap();
499 conn.commit().await.unwrap();
500
501 let rows = conn.query("SELECT * FROM tx_test").await.unwrap();
502 assert_eq!(rows.len(), 2);
503 }
504
505 #[tokio::test]
506 async fn test_any_pool_sqlite_transaction_rollback() {
507 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
508 let mut conn = pool.create().await.unwrap();
509
510 conn.execute("CREATE TABLE tx_rb (id INTEGER PRIMARY KEY, val INTEGER)")
511 .await
512 .unwrap();
513
514 conn.begin_transaction().await.unwrap();
515 conn.execute("INSERT INTO tx_rb (val) VALUES (1)")
516 .await
517 .unwrap();
518 conn.rollback().await.unwrap();
519
520 let rows = conn.query("SELECT * FROM tx_rb").await.unwrap();
521 assert_eq!(rows.len(), 0);
522 }
523
524 #[tokio::test]
525 async fn test_any_pool_sqlite_ping() {
526 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
527 let mut conn = pool.create().await.unwrap();
528 let ok = conn.ping().await;
529 assert!(ok);
530 assert!(conn.is_connected());
531 }
532
533 #[tokio::test]
534 async fn test_any_pool_invalid_dsn() {
535 let result = AnyPool::connect("invalid://dsn").await;
536 assert!(result.is_err());
537 }
538
539 #[tokio::test]
540 async fn test_any_pool_sqlite_count_query() {
541 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
542 let mut conn = pool.create().await.unwrap();
543
544 conn.execute("CREATE TABLE cnt (id INTEGER PRIMARY KEY)")
545 .await
546 .unwrap();
547 for i in 1..=5 {
548 conn.execute(&format!("INSERT INTO cnt (id) VALUES ({})", i))
549 .await
550 .unwrap();
551 }
552
553 let rows = conn.query("SELECT * FROM cnt").await.unwrap();
555 assert_eq!(rows.len(), 5);
556 }
557
558 #[test]
561 fn test_any_backend_dialect_mapping() {
562 use sz_orm_core::DbType;
563
564 let mysql_d = AnyBackend::MySql.dialect();
565 assert_eq!(mysql_d.db_type(), DbType::MySQL);
566
567 let pg_d = AnyBackend::Postgres.dialect();
568 assert_eq!(pg_d.db_type(), DbType::PostgreSQL);
569
570 let sqlite_d = AnyBackend::Sqlite.dialect();
571 assert_eq!(sqlite_d.db_type(), DbType::Sqlite);
572
573 let oracle_d = AnyBackend::Oracle.dialect();
574 assert_eq!(oracle_d.db_type(), DbType::Oracle);
575
576 let mssql_d = AnyBackend::Mssql.dialect();
577 assert_eq!(mssql_d.db_type(), DbType::SqlServer);
578 }
579
580 #[test]
581 fn test_oracle_dialect_pagination() {
582 let d = AnyBackend::Oracle.dialect();
583 let sql = d.build_pagination("SELECT * FROM users", 2, 10);
584 let upper = sql.to_uppercase();
585 assert!(
586 upper.contains("OFFSET") || upper.contains("FETCH") || upper.contains("ROWNUM"),
587 "Oracle 分页 SQL 应含 OFFSET/FETCH/ROWNUM,实际: {}",
588 sql
589 );
590 assert!(
591 !upper.contains("LIMIT"),
592 "Oracle 分页 SQL 不应含 LIMIT,实际: {}",
593 sql
594 );
595 }
596
597 #[test]
598 fn test_mssql_dialect_pagination() {
599 let d = AnyBackend::Mssql.dialect();
600 let sql = d.build_pagination("SELECT * FROM users", 2, 10);
601 let upper = sql.to_uppercase();
602 assert!(
603 upper.contains("OFFSET") || upper.contains("FETCH"),
604 "MSSQL 分页 SQL 应含 OFFSET/FETCH,实际: {}",
605 sql
606 );
607 assert!(
608 !upper.contains("LIMIT"),
609 "MSSQL 分页 SQL 不应含 LIMIT,实际: {}",
610 sql
611 );
612 }
613
614 #[test]
615 fn test_oracle_dialect_no_placeholder() {
616 let d = AnyBackend::Oracle.dialect();
617 let _ = d.db_type();
619 let _ = d.quote("col");
620 let _ = d.quote_checked("col").unwrap();
621 let _ = d.escape_string("val");
622 let _ = d.supports_returning();
623 let _ = d.build_pagination("SELECT 1", 1, 10);
624 let _ = d.json_type();
625 let _ = d.json_extract("col", "$.key");
626 let _ = d.full_text_search(&["col"], "kw");
627 let _ = d.bool_to_int("expr");
628 let _ = d.concat(&["a", "b"]);
629 }
630
631 #[test]
632 fn test_mssql_dialect_no_placeholder() {
633 let d = AnyBackend::Mssql.dialect();
634 let _ = d.db_type();
635 let _ = d.quote("col");
636 let _ = d.quote_checked("col").unwrap();
637 let _ = d.escape_string("val");
638 let _ = d.supports_returning();
639 let _ = d.build_pagination("SELECT 1", 1, 10);
640 let _ = d.json_type();
641 let _ = d.json_extract("col", "$.key");
642 let _ = d.full_text_search(&["col"], "kw");
643 let _ = d.bool_to_int("expr");
644 let _ = d.concat(&["a", "b"]);
645 }
646
647 #[tokio::test]
648 async fn test_any_pool_dialect_sqlite() {
649 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
650 let d = pool.dialect();
651 assert_eq!(d.db_type(), sz_orm_core::DbType::Sqlite);
652 }
653}