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 sz_orm_core::DbType;
40
41use crate::any::{
42 MySqlPoolHandle, PgPoolHandle, SqlitePoolHandle, SqlxMySqlConnectionFactory,
43 SqlxPgConnectionFactory, SqlxSqliteConnectionFactory,
44};
45
46#[cfg(feature = "oracle")]
47use sz_orm_oracle::{OracleConnectionFactory, OraclePoolHandle};
48
49#[cfg(feature = "mssql")]
50use sz_orm_mssql::{MssqlConnectionFactory, MssqlPoolHandle};
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58#[non_exhaustive]
59pub enum AnyBackend {
60 MySql,
62 Postgres,
64 Sqlite,
66 Oracle,
68 Mssql,
70}
71
72impl AnyBackend {
73 pub fn from_dsn(dsn: &str) -> Result<Self, DbError> {
83 if dsn.starts_with("mysql://") || dsn.starts_with("mariadb://") {
84 Ok(AnyBackend::MySql)
85 } else if dsn.starts_with("postgres://") || dsn.starts_with("postgresql://") {
86 Ok(AnyBackend::Postgres)
87 } else if dsn.starts_with("sqlite://") || dsn.starts_with("sqlite:") {
88 Ok(AnyBackend::Sqlite)
89 } else if dsn.starts_with("oracle://") {
90 Ok(AnyBackend::Oracle)
91 } else if dsn.starts_with("mssql://") || dsn.starts_with("sqlserver://") {
92 Ok(AnyBackend::Mssql)
93 } else {
94 Err(DbError::ConnectionRefused(format!(
95 "未知的 DSN scheme: {}(支持 mysql/postgres/sqlite/oracle/mssql)",
96 dsn
97 )))
98 }
99 }
100
101 pub fn from_db_type(db_type: DbType) -> Option<Self> {
115 match db_type {
116 DbType::MySQL | DbType::MariaDB | DbType::TiDB | DbType::OceanBase => {
117 Some(AnyBackend::MySql)
118 }
119 DbType::PostgreSQL | DbType::Kingbase | DbType::PolarDB | DbType::GaussDB => {
120 Some(AnyBackend::Postgres)
121 }
122 DbType::Sqlite => Some(AnyBackend::Sqlite),
123 DbType::Oracle | DbType::Dameng => Some(AnyBackend::Oracle),
124 DbType::SqlServer | DbType::Sybase | DbType::GBase => Some(AnyBackend::Mssql),
125 _ => None,
126 }
127 }
128
129 pub fn name(&self) -> &'static str {
131 match self {
132 AnyBackend::MySql => "mysql",
133 AnyBackend::Postgres => "postgres",
134 AnyBackend::Sqlite => "sqlite",
135 AnyBackend::Oracle => "oracle",
136 AnyBackend::Mssql => "mssql",
137 }
138 }
139
140 pub fn dialect(&self) -> Box<dyn Dialect> {
148 match self {
149 AnyBackend::MySql => Box::new(MySqlDialect),
150 AnyBackend::Postgres => Box::new(PostgreSqlDialect),
151 AnyBackend::Sqlite => Box::new(SqliteDialect),
152 AnyBackend::Oracle => Box::new(OracleDialect),
153 AnyBackend::Mssql => Box::new(SqlServerDialect),
154 }
155 }
156}
157
158pub struct AnyPool {
160 backend: AnyBackend,
161 factory: Arc<dyn ConnectionFactory>,
162}
163
164impl AnyPool {
165 pub async fn connect(dsn: &str) -> Result<Self, DbError> {
173 let backend = AnyBackend::from_dsn(dsn)?;
174 let factory: Arc<dyn ConnectionFactory> = match backend {
175 AnyBackend::MySql => {
176 let handle = Arc::new(MySqlPoolHandle::connect(dsn).await?);
177 Arc::new(SqlxMySqlConnectionFactory::new(handle))
178 }
179 AnyBackend::Postgres => {
180 let handle = Arc::new(PgPoolHandle::connect(dsn).await?);
181 Arc::new(SqlxPgConnectionFactory::new(handle))
182 }
183 AnyBackend::Sqlite => {
184 let handle = Arc::new(SqlitePoolHandle::connect(dsn).await?);
185 Arc::new(SqlxSqliteConnectionFactory::new(handle))
186 }
187 AnyBackend::Oracle => {
188 #[cfg(feature = "oracle")]
189 {
190 let (username, password, connect_string) = parse_oracle_dsn(dsn)?;
191 let handle = Arc::new(OraclePoolHandle::connect(
192 &username,
193 &password,
194 &connect_string,
195 )?);
196 Arc::new(OracleConnectionFactory::new(handle))
197 }
198 #[cfg(not(feature = "oracle"))]
199 {
200 return Err(DbError::ConnectionRefused(
201 "Oracle 后端未启用,请在 Cargo.toml 中添加 features = [\"oracle\"]"
202 .to_string(),
203 ));
204 }
205 }
206 AnyBackend::Mssql => {
207 #[cfg(feature = "mssql")]
208 {
209 let ado_string = parse_mssql_dsn(dsn)?;
210 let handle = Arc::new(MssqlPoolHandle::connect(&ado_string).await?);
211 Arc::new(MssqlConnectionFactory::new(handle))
212 }
213 #[cfg(not(feature = "mssql"))]
214 {
215 return Err(DbError::ConnectionRefused(
216 "MSSQL 后端未启用,请在 Cargo.toml 中添加 features = [\"mssql\"]"
217 .to_string(),
218 ));
219 }
220 }
221 };
222 Ok(Self { backend, factory })
223 }
224
225 pub fn from_factory(backend: AnyBackend, factory: Arc<dyn ConnectionFactory>) -> Self {
227 Self { backend, factory }
228 }
229
230 pub fn backend(&self) -> AnyBackend {
232 self.backend
233 }
234
235 pub fn dialect(&self) -> Box<dyn Dialect> {
239 self.backend.dialect()
240 }
241
242 pub async fn create(&self) -> Result<AnyConnection, DbError> {
244 let conn = self.factory.create().await?;
245 Ok(AnyConnection {
246 backend: self.backend,
247 inner: conn,
248 })
249 }
250}
251
252pub async fn create_connection(dsn: &str) -> Result<Box<dyn Connection>, DbError> {
283 let pool = AnyPool::connect(dsn).await?;
284 let conn = pool.create().await?;
285 Ok(Box::new(conn))
286}
287
288pub async fn create_connection_by_type(
307 db_type: DbType,
308 dsn: &str,
309) -> Result<Box<dyn Connection>, DbError> {
310 let backend = AnyBackend::from_db_type(db_type).ok_or_else(|| {
311 DbError::ConnectionRefused(format!(
312 "DbType {:?} 不在 sz-orm-sqlx 支持范围内(支持 MySQL/PostgreSQL/SQLite/Oracle/MSSQL 及兼容方言)",
313 db_type
314 ))
315 })?;
316 let normalized = normalize_dsn_scheme(backend, dsn);
317 create_connection(&normalized).await
318}
319
320fn normalize_dsn_scheme(backend: AnyBackend, dsn: &str) -> String {
325 let expected_scheme = match backend {
326 AnyBackend::MySql => "mysql://",
327 AnyBackend::Postgres => "postgres://",
328 AnyBackend::Sqlite => "sqlite:",
329 AnyBackend::Oracle => "oracle://",
330 AnyBackend::Mssql => "mssql://",
331 };
332 let known_schemes = [
333 "mysql://",
334 "mariadb://",
335 "postgres://",
336 "postgresql://",
337 "sqlite://",
338 "sqlite:",
339 "oracle://",
340 "mssql://",
341 "sqlserver://",
342 ];
343 for scheme in &known_schemes {
344 if let Some(rest) = dsn.strip_prefix(scheme) {
345 if backend == AnyBackend::from_dsn(dsn).unwrap_or(AnyBackend::MySql) {
346 return dsn.to_string();
347 }
348 return format!("{}{}", expected_scheme, rest);
349 }
350 }
351 dsn.to_string()
352}
353
354#[allow(dead_code)]
360pub(crate) fn parse_oracle_dsn(dsn: &str) -> Result<(String, String, String), DbError> {
361 let rest = dsn
362 .strip_prefix("oracle://")
363 .ok_or_else(|| DbError::ConnectionRefused(format!("无效的 Oracle DSN: {}", dsn)))?;
364 parse_user_pass_host(rest, "Oracle")
365}
366
367#[allow(dead_code)]
373pub(crate) fn parse_mssql_dsn(dsn: &str) -> Result<String, DbError> {
374 let rest = dsn
375 .strip_prefix("mssql://")
376 .or_else(|| dsn.strip_prefix("sqlserver://"))
377 .ok_or_else(|| DbError::ConnectionRefused(format!("无效的 MSSQL DSN: {}", dsn)))?;
378 let (username, password, host_port_db) = parse_user_pass_host(rest, "MSSQL")?;
379 let (host_port, database) = host_port_db
380 .split_once('/')
381 .ok_or_else(|| DbError::ConnectionRefused(format!("MSSQL DSN 缺少 database: {}", dsn)))?;
382 let (host, port) = host_port.split_once(':').unwrap_or((host_port, "1433"));
383 Ok(format!(
384 "Server={},{};Database={};User Id={};Password={};",
385 host, port, database, username, password
386 ))
387}
388
389#[allow(dead_code)]
391fn parse_user_pass_host(rest: &str, backend: &str) -> Result<(String, String, String), DbError> {
392 let (userinfo, hostinfo) = rest
393 .split_once('@')
394 .ok_or_else(|| DbError::ConnectionRefused(format!("{} DSN 缺少 @: {}", backend, rest)))?;
395 let (username, password) = userinfo.split_once(':').ok_or_else(|| {
396 DbError::ConnectionRefused(format!("{} DSN 缺少 password: {}", backend, rest))
397 })?;
398 Ok((
399 username.to_string(),
400 password.to_string(),
401 hostinfo.to_string(),
402 ))
403}
404
405pub struct AnyConnection {
407 backend: AnyBackend,
408 inner: Box<dyn Connection>,
409}
410
411impl AnyConnection {
412 pub fn backend(&self) -> AnyBackend {
414 self.backend
415 }
416}
417
418impl Connection for AnyConnection {
419 fn execute<'a>(
420 &'a mut self,
421 sql: &'a str,
422 ) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
423 self.inner.execute(sql)
424 }
425
426 fn query<'a>(
427 &'a mut self,
428 sql: &'a str,
429 ) -> Pin<Box<dyn Future<Output = Result<QueryRows, DbError>> + Send + 'a>> {
430 self.inner.query(sql)
431 }
432
433 fn begin_transaction<'a>(
434 &'a mut self,
435 ) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
436 self.inner.begin_transaction()
437 }
438
439 fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
440 self.inner.commit()
441 }
442
443 fn rollback<'a>(
444 &'a mut self,
445 ) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
446 self.inner.rollback()
447 }
448
449 fn is_connected(&self) -> bool {
450 self.inner.is_connected()
451 }
452
453 fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
454 self.inner.ping()
455 }
456
457 fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
458 self.inner.close()
459 }
460}
461
462#[cfg(test)]
463mod tests {
464 use super::*;
465
466 #[test]
467 fn test_any_backend_from_dsn_mysql() {
468 assert_eq!(
469 AnyBackend::from_dsn("mysql://root:pass@127.0.0.1/db").unwrap(),
470 AnyBackend::MySql
471 );
472 assert_eq!(
473 AnyBackend::from_dsn("mariadb://root:pass@127.0.0.1/db").unwrap(),
474 AnyBackend::MySql
475 );
476 }
477
478 #[test]
479 fn test_any_backend_from_dsn_postgres() {
480 assert_eq!(
481 AnyBackend::from_dsn("postgres://user:pass@127.0.0.1/db").unwrap(),
482 AnyBackend::Postgres
483 );
484 assert_eq!(
485 AnyBackend::from_dsn("postgresql://user:pass@127.0.0.1/db").unwrap(),
486 AnyBackend::Postgres
487 );
488 }
489
490 #[test]
491 fn test_any_backend_from_dsn_sqlite() {
492 assert_eq!(
493 AnyBackend::from_dsn("sqlite::memory:").unwrap(),
494 AnyBackend::Sqlite
495 );
496 assert_eq!(
497 AnyBackend::from_dsn("sqlite://./test.db").unwrap(),
498 AnyBackend::Sqlite
499 );
500 }
501
502 #[test]
503 fn test_any_backend_from_dsn_oracle() {
504 assert_eq!(
505 AnyBackend::from_dsn("oracle://sys:test123@127.0.0.1:1521/freepdb1").unwrap(),
506 AnyBackend::Oracle
507 );
508 }
509
510 #[test]
511 fn test_any_backend_from_dsn_mssql() {
512 assert_eq!(
513 AnyBackend::from_dsn("mssql://sa:test123@127.0.0.1:1433/testdb").unwrap(),
514 AnyBackend::Mssql
515 );
516 }
517
518 #[test]
519 fn test_any_backend_from_dsn_sqlserver() {
520 assert_eq!(
521 AnyBackend::from_dsn("sqlserver://sa:test123@127.0.0.1:1433/testdb").unwrap(),
522 AnyBackend::Mssql
523 );
524 }
525
526 #[test]
527 fn test_any_backend_from_dsn_unknown_v22() {
528 let result = AnyBackend::from_dsn("redis://127.0.0.1");
529 assert!(result.is_err());
530 if let Err(e) = result {
531 let msg = format!("{}", e);
532 assert!(msg.contains("mysql"));
533 assert!(msg.contains("postgres"));
534 assert!(msg.contains("sqlite"));
535 assert!(msg.contains("oracle"));
536 assert!(msg.contains("mssql"));
537 }
538 }
539
540 #[test]
541 fn test_any_backend_name_v22() {
542 assert_eq!(AnyBackend::MySql.name(), "mysql");
543 assert_eq!(AnyBackend::Postgres.name(), "postgres");
544 assert_eq!(AnyBackend::Sqlite.name(), "sqlite");
545 assert_eq!(AnyBackend::Oracle.name(), "oracle");
546 assert_eq!(AnyBackend::Mssql.name(), "mssql");
547 }
548
549 #[test]
550 fn test_any_backend_equality() {
551 assert_eq!(AnyBackend::MySql, AnyBackend::MySql);
552 assert_ne!(AnyBackend::MySql, AnyBackend::Postgres);
553 assert_ne!(AnyBackend::Postgres, AnyBackend::Sqlite);
554 assert_ne!(AnyBackend::Oracle, AnyBackend::Mssql);
555 assert_ne!(AnyBackend::Oracle, AnyBackend::MySql);
556 }
557
558 #[test]
559 fn test_parse_oracle_dsn() {
560 let (user, pass, cs) =
561 parse_oracle_dsn("oracle://sys:test123@127.0.0.1:1521/freepdb1").unwrap();
562 assert_eq!(user, "sys");
563 assert_eq!(pass, "test123");
564 assert_eq!(cs, "127.0.0.1:1521/freepdb1");
565 }
566
567 #[test]
568 fn test_parse_mssql_dsn() {
569 let ado = parse_mssql_dsn("mssql://sa:test123@127.0.0.1:1433/testdb").unwrap();
570 assert!(ado.contains("Server=127.0.0.1,1433"));
571 assert!(ado.contains("Database=testdb"));
572 assert!(ado.contains("User Id=sa"));
573 assert!(ado.contains("Password=test123"));
574 }
575
576 #[test]
577 fn test_parse_mssql_dsn_sqlserver_scheme() {
578 let ado = parse_mssql_dsn("sqlserver://sa:pass@localhost/testdb").unwrap();
579 assert!(ado.contains("Server=localhost,1433"));
580 assert!(ado.contains("Database=testdb"));
581 }
582
583 #[tokio::test]
586 async fn test_any_pool_sqlite_connect_and_query() {
587 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
588 assert_eq!(pool.backend(), AnyBackend::Sqlite);
589
590 let mut conn = pool.create().await.unwrap();
591 assert_eq!(conn.backend(), AnyBackend::Sqlite);
592
593 conn.execute("CREATE TABLE test_any (id INTEGER PRIMARY KEY, name TEXT)")
595 .await
596 .unwrap();
597 conn.execute("INSERT INTO test_any (name) VALUES ('Alice')")
598 .await
599 .unwrap();
600 conn.execute("INSERT INTO test_any (name) VALUES ('Bob')")
601 .await
602 .unwrap();
603
604 let rows = conn
606 .query("SELECT * FROM test_any ORDER BY id")
607 .await
608 .unwrap();
609 assert_eq!(rows.len(), 2);
610 assert_eq!(rows[0].get("name").and_then(|v| v.as_str()), Some("Alice"));
611 assert_eq!(rows[1].get("name").and_then(|v| v.as_str()), Some("Bob"));
612 }
613
614 #[tokio::test]
615 async fn test_any_pool_sqlite_transaction_commit() {
616 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
617 let mut conn = pool.create().await.unwrap();
618
619 conn.execute("CREATE TABLE tx_test (id INTEGER PRIMARY KEY, val INTEGER)")
620 .await
621 .unwrap();
622
623 conn.begin_transaction().await.unwrap();
625 conn.execute("INSERT INTO tx_test (val) VALUES (1)")
626 .await
627 .unwrap();
628 conn.execute("INSERT INTO tx_test (val) VALUES (2)")
629 .await
630 .unwrap();
631 conn.commit().await.unwrap();
632
633 let rows = conn.query("SELECT * FROM tx_test").await.unwrap();
634 assert_eq!(rows.len(), 2);
635 }
636
637 #[tokio::test]
638 async fn test_any_pool_sqlite_transaction_rollback() {
639 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
640 let mut conn = pool.create().await.unwrap();
641
642 conn.execute("CREATE TABLE tx_rb (id INTEGER PRIMARY KEY, val INTEGER)")
643 .await
644 .unwrap();
645
646 conn.begin_transaction().await.unwrap();
647 conn.execute("INSERT INTO tx_rb (val) VALUES (1)")
648 .await
649 .unwrap();
650 conn.rollback().await.unwrap();
651
652 let rows = conn.query("SELECT * FROM tx_rb").await.unwrap();
653 assert_eq!(rows.len(), 0);
654 }
655
656 #[tokio::test]
657 async fn test_any_pool_sqlite_ping() {
658 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
659 let mut conn = pool.create().await.unwrap();
660 let ok = conn.ping().await;
661 assert!(ok);
662 assert!(conn.is_connected());
663 }
664
665 #[tokio::test]
666 async fn test_any_pool_invalid_dsn() {
667 let result = AnyPool::connect("invalid://dsn").await;
668 assert!(result.is_err());
669 }
670
671 #[tokio::test]
672 async fn test_any_pool_sqlite_count_query() {
673 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
674 let mut conn = pool.create().await.unwrap();
675
676 conn.execute("CREATE TABLE cnt (id INTEGER PRIMARY KEY)")
677 .await
678 .unwrap();
679 for i in 1..=5 {
680 conn.execute(&format!("INSERT INTO cnt (id) VALUES ({})", i))
681 .await
682 .unwrap();
683 }
684
685 let rows = conn.query("SELECT * FROM cnt").await.unwrap();
687 assert_eq!(rows.len(), 5);
688 }
689
690 #[test]
693 fn test_any_backend_dialect_mapping() {
694 use sz_orm_core::DbType;
695
696 let mysql_d = AnyBackend::MySql.dialect();
697 assert_eq!(mysql_d.db_type(), DbType::MySQL);
698
699 let pg_d = AnyBackend::Postgres.dialect();
700 assert_eq!(pg_d.db_type(), DbType::PostgreSQL);
701
702 let sqlite_d = AnyBackend::Sqlite.dialect();
703 assert_eq!(sqlite_d.db_type(), DbType::Sqlite);
704
705 let oracle_d = AnyBackend::Oracle.dialect();
706 assert_eq!(oracle_d.db_type(), DbType::Oracle);
707
708 let mssql_d = AnyBackend::Mssql.dialect();
709 assert_eq!(mssql_d.db_type(), DbType::SqlServer);
710 }
711
712 #[test]
713 fn test_oracle_dialect_pagination() {
714 let d = AnyBackend::Oracle.dialect();
715 let sql = d.build_pagination("SELECT * FROM users", 2, 10);
716 let upper = sql.to_uppercase();
717 assert!(
718 upper.contains("OFFSET") || upper.contains("FETCH") || upper.contains("ROWNUM"),
719 "Oracle 分页 SQL 应含 OFFSET/FETCH/ROWNUM,实际: {}",
720 sql
721 );
722 assert!(
723 !upper.contains("LIMIT"),
724 "Oracle 分页 SQL 不应含 LIMIT,实际: {}",
725 sql
726 );
727 }
728
729 #[test]
730 fn test_mssql_dialect_pagination() {
731 let d = AnyBackend::Mssql.dialect();
732 let sql = d.build_pagination("SELECT * FROM users", 2, 10);
733 let upper = sql.to_uppercase();
734 assert!(
735 upper.contains("OFFSET") || upper.contains("FETCH"),
736 "MSSQL 分页 SQL 应含 OFFSET/FETCH,实际: {}",
737 sql
738 );
739 assert!(
740 !upper.contains("LIMIT"),
741 "MSSQL 分页 SQL 不应含 LIMIT,实际: {}",
742 sql
743 );
744 }
745
746 #[test]
747 fn test_oracle_dialect_no_placeholder() {
748 let d = AnyBackend::Oracle.dialect();
749 let _ = d.db_type();
751 let _ = d.quote("col");
752 let _ = d.quote_checked("col").unwrap();
753 let _ = d.escape_string("val");
754 let _ = d.supports_returning();
755 let _ = d.build_pagination("SELECT 1", 1, 10);
756 let _ = d.json_type();
757 let _ = d.json_extract("col", "$.key");
758 let _ = d.full_text_search(&["col"], "kw");
759 let _ = d.bool_to_int("expr");
760 let _ = d.concat(&["a", "b"]);
761 }
762
763 #[test]
764 fn test_mssql_dialect_no_placeholder() {
765 let d = AnyBackend::Mssql.dialect();
766 let _ = d.db_type();
767 let _ = d.quote("col");
768 let _ = d.quote_checked("col").unwrap();
769 let _ = d.escape_string("val");
770 let _ = d.supports_returning();
771 let _ = d.build_pagination("SELECT 1", 1, 10);
772 let _ = d.json_type();
773 let _ = d.json_extract("col", "$.key");
774 let _ = d.full_text_search(&["col"], "kw");
775 let _ = d.bool_to_int("expr");
776 let _ = d.concat(&["a", "b"]);
777 }
778
779 #[tokio::test]
780 async fn test_any_pool_dialect_sqlite() {
781 let pool = AnyPool::connect("sqlite::memory:").await.unwrap();
782 let d = pool.dialect();
783 assert_eq!(d.db_type(), sz_orm_core::DbType::Sqlite);
784 }
785
786 #[test]
789 fn test_from_db_type_supported() {
790 use sz_orm_core::DbType;
791 assert_eq!(
792 AnyBackend::from_db_type(DbType::MySQL),
793 Some(AnyBackend::MySql)
794 );
795 assert_eq!(
796 AnyBackend::from_db_type(DbType::MariaDB),
797 Some(AnyBackend::MySql)
798 );
799 assert_eq!(
800 AnyBackend::from_db_type(DbType::TiDB),
801 Some(AnyBackend::MySql)
802 );
803 assert_eq!(
804 AnyBackend::from_db_type(DbType::OceanBase),
805 Some(AnyBackend::MySql)
806 );
807 assert_eq!(
808 AnyBackend::from_db_type(DbType::PostgreSQL),
809 Some(AnyBackend::Postgres)
810 );
811 assert_eq!(
812 AnyBackend::from_db_type(DbType::Kingbase),
813 Some(AnyBackend::Postgres)
814 );
815 assert_eq!(
816 AnyBackend::from_db_type(DbType::PolarDB),
817 Some(AnyBackend::Postgres)
818 );
819 assert_eq!(
820 AnyBackend::from_db_type(DbType::GaussDB),
821 Some(AnyBackend::Postgres)
822 );
823 assert_eq!(
824 AnyBackend::from_db_type(DbType::Sqlite),
825 Some(AnyBackend::Sqlite)
826 );
827 assert_eq!(
828 AnyBackend::from_db_type(DbType::Oracle),
829 Some(AnyBackend::Oracle)
830 );
831 assert_eq!(
832 AnyBackend::from_db_type(DbType::Dameng),
833 Some(AnyBackend::Oracle)
834 );
835 assert_eq!(
836 AnyBackend::from_db_type(DbType::SqlServer),
837 Some(AnyBackend::Mssql)
838 );
839 assert_eq!(
840 AnyBackend::from_db_type(DbType::Sybase),
841 Some(AnyBackend::Mssql)
842 );
843 assert_eq!(
844 AnyBackend::from_db_type(DbType::GBase),
845 Some(AnyBackend::Mssql)
846 );
847 }
848
849 #[test]
850 fn test_from_db_type_unsupported() {
851 use sz_orm_core::DbType;
852 assert_eq!(AnyBackend::from_db_type(DbType::Redis), None);
853 assert_eq!(AnyBackend::from_db_type(DbType::MongoDB), None);
854 assert_eq!(AnyBackend::from_db_type(DbType::ClickHouse), None);
855 assert_eq!(AnyBackend::from_db_type(DbType::VectorDb), None);
856 }
857
858 #[tokio::test]
859 async fn test_create_connection_sqlite() {
860 use sz_orm_core::Connection;
861 let mut conn = create_connection("sqlite::memory:").await.unwrap();
862 conn.execute("CREATE TABLE cc (id INTEGER PRIMARY KEY, v TEXT)")
863 .await
864 .unwrap();
865 conn.execute("INSERT INTO cc (v) VALUES ('ok')")
866 .await
867 .unwrap();
868 let rows = conn.query("SELECT * FROM cc").await.unwrap();
869 assert_eq!(rows.len(), 1);
870 }
871
872 #[tokio::test]
873 async fn test_create_connection_invalid_dsn() {
874 let result = create_connection("invalid://dsn").await;
875 assert!(result.is_err());
876 }
877
878 #[tokio::test]
879 async fn test_create_connection_by_type_sqlite() {
880 use sz_orm_core::{Connection, DbType};
881 let mut conn = create_connection_by_type(DbType::Sqlite, "sqlite::memory:")
882 .await
883 .unwrap();
884 conn.execute("SELECT 1").await.unwrap();
885 }
886
887 #[tokio::test]
888 async fn test_create_connection_by_type_unsupported() {
889 use sz_orm_core::DbType;
890 let result = create_connection_by_type(DbType::Redis, "redis://127.0.0.1").await;
891 assert!(result.is_err());
892 }
893
894 #[test]
895 fn test_normalize_dsn_scheme_already_correct() {
896 assert_eq!(
897 normalize_dsn_scheme(AnyBackend::MySql, "mysql://root:pass@host/db"),
898 "mysql://root:pass@host/db"
899 );
900 assert_eq!(
901 normalize_dsn_scheme(AnyBackend::Postgres, "postgres://user@host/db"),
902 "postgres://user@host/db"
903 );
904 assert_eq!(
905 normalize_dsn_scheme(AnyBackend::Oracle, "oracle://sys:pass@host/svc"),
906 "oracle://sys:pass@host/svc"
907 );
908 }
909
910 #[test]
911 fn test_normalize_dsn_scheme_fix_mismatch() {
912 let fixed = normalize_dsn_scheme(AnyBackend::MySql, "postgres://root:pass@host/db");
913 assert!(fixed.starts_with("mysql://"));
914 assert!(fixed.contains("root:pass@host/db"));
915 }
916}