Skip to main content

sz_orm_sqlx/
any_driver.rs

1//! Any driver — 一套代码多 DB 后端透明切换(SQLx 风格)
2//!
3//! SQLx 提供 `sqlx::Any` 适配器,让同一份代码可以在 MySQL/PostgreSQL/SQLite
4//! 之间透明切换。SZ-ORM 在 `sz-orm-sqlx` 已有各后端独立实现,
5//! 此模块在上层提供统一的 [`AnyConnection`] 和 [`AnyPool`] 抽象,
6//! 让运行时切换数据库后端成为可能。
7//!
8//! # 设计
9//!
10//! - [`AnyBackend`]:枚举后端类型
11//! - [`AnyPool`]:持有具体后端的 `Box<dyn ConnectionFactory>`
12//! - [`AnyConnection`]:持有具体后端的 `Box<dyn Connection>`
13//! - 通过 DSN 自动识别后端类型,运行时透明切换
14//!
15//! # 用法
16//!
17//! ```ignore
18//! use sz_orm_sqlx::any_driver::{AnyBackend, AnyPool};
19//!
20//! // 从 DSN 自动识别后端
21//! let pool = AnyPool::connect("mysql://root:pass@127.0.0.1/db").await?;
22//! let mut conn = pool.create().await?;
23//! let rows = conn.query("SELECT 1").await?;
24//!
25//! // 运行时切换后端
26//! let pg_pool = AnyPool::connect("postgres://user:pass@127.0.0.1/db").await?;
27//! let mut pg_conn = pg_pool.create().await?;
28//! let rows = pg_conn.query("SELECT 1").await?;
29//! ```
30
31use 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/// 数据库后端类型
53///
54/// v2.2.0 新增 `Oracle` 和 `Mssql` 变体(需启用 `oracle`/`mssql` feature)。
55/// `#[non_exhaustive]` 标注确保外部 crate match 时必须使用 wildcard 臂,
56/// 未来新增变体不会破坏现有代码。
57#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58#[non_exhaustive]
59pub enum AnyBackend {
60    /// MySQL / MariaDB
61    MySql,
62    /// PostgreSQL
63    Postgres,
64    /// SQLite
65    Sqlite,
66    /// Oracle(v2.2.0 新增,需启用 `oracle` feature)
67    Oracle,
68    /// SQL Server / MSSQL(v2.2.0 新增,需启用 `mssql` feature)
69    Mssql,
70}
71
72impl AnyBackend {
73    /// 从 DSN 自动识别后端类型
74    ///
75    /// # 支持的 scheme
76    ///
77    /// - `mysql://` / `mariadb://` → MySQL
78    /// - `postgres://` / `postgresql://` → Postgres
79    /// - `sqlite://` / `sqlite:` → Sqlite
80    /// - `oracle://` → Oracle(v2.2.0 新增)
81    /// - `mssql://` / `sqlserver://` → Mssql(v2.2.0 新增)
82    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    /// 从 [`DbType`] 转换为 [`AnyBackend`]
102    ///
103    /// 返回 `None` 的情况:DbType 对应的后端不在 sz-orm-sqlx 支持范围内
104    /// (如 Redis、MongoDB、ClickHouse 等非关系型数据库)。
105    ///
106    /// # 映射
107    ///
108    /// - `DbType::MySQL` / `DbType::MariaDB` / `DbType::TiDB` / `DbType::OceanBase` → MySql
109    /// - `DbType::PostgreSQL` / `DbType::Kingbase` / `DbType::PolarDB` / `DbType::GaussDB` → Postgres
110    /// - `DbType::Sqlite` → Sqlite
111    /// - `DbType::Oracle` / `DbType::Dameng` → Oracle
112    /// - `DbType::SqlServer` / `DbType::Sybase` / `DbType::GBase` → Mssql
113    /// - 其他 → None
114    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    /// 后端名称
130    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    /// 返回对应后端的 Dialect 实例(v2.2.0 新增)
141    ///
142    /// - MySql → [`MySqlDialect`]
143    /// - Postgres → [`PostgreSqlDialect`]
144    /// - Sqlite → [`SqliteDialect`]
145    /// - Oracle → [`OracleDialect`]
146    /// - Mssql → [`SqlServerDialect`]
147    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
158/// 后端无关的连接工厂
159pub struct AnyPool {
160    backend: AnyBackend,
161    factory: Arc<dyn ConnectionFactory>,
162}
163
164impl AnyPool {
165    /// 连接数据库,根据 DSN 自动识别后端
166    ///
167    /// # 错误
168    ///
169    /// - DSN scheme 不识别 → [`DbError::ConnectionRefused`]
170    /// - 连接失败 → [`DbError::ConnectionError`]
171    /// - Oracle/MSSQL 后端未启用对应 feature → [`DbError::ConnectionRefused`] 含提示
172    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    /// 从已有的连接工厂构造
226    pub fn from_factory(backend: AnyBackend, factory: Arc<dyn ConnectionFactory>) -> Self {
227        Self { backend, factory }
228    }
229
230    /// 获取后端类型
231    pub fn backend(&self) -> AnyBackend {
232        self.backend
233    }
234
235    /// 返回对应后端的 Dialect 实例(v2.2.0 新增)
236    ///
237    /// 委托 [`AnyBackend::dialect()`],根据后端自动选择方言。
238    pub fn dialect(&self) -> Box<dyn Dialect> {
239        self.backend.dialect()
240    }
241
242    /// 创建一个新连接
243    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
252/// 统一连接便利函数:按 DSN scheme 自动识别后端并创建连接
253///
254/// 这是 sz-orm-sqlx 对外提供的统一入口,下游项目(如 sz-rust)只需调用此函数,
255/// 无需为每种数据库编写独立的连接逻辑。Oracle/MSSQL 支持通过 feature gate 启用。
256///
257/// # 支持的 DSN scheme
258///
259/// - `mysql://` / `mariadb://` → MySQL
260/// - `postgres://` / `postgresql://` → PostgreSQL
261/// - `sqlite://` / `sqlite:` → SQLite
262/// - `oracle://` → Oracle(需启用 `oracle` feature)
263/// - `mssql://` / `sqlserver://` → MSSQL(需启用 `mssql` feature)
264///
265/// # 错误
266///
267/// - DSN scheme 不识别 → [`DbError::ConnectionRefused`]
268/// - Oracle/MSSQL 后端未启用对应 feature → [`DbError::ConnectionRefused`] 含提示
269/// - 连接失败 → [`DbError::ConnectionError`]
270///
271/// # 示例
272///
273/// ```no_run
274/// # #[tokio::main]
275/// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
276/// use sz_orm_core::Connection;
277/// let mut conn = sz_orm_sqlx::create_connection("sqlite::memory:").await?;
278/// conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY)").await?;
279/// # Ok(())
280/// # }
281/// ```
282pub 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
288/// 按 [`DbType`] 创建连接
289///
290/// 当 DSN 的 scheme 可能与 `db_type` 不一致时使用此函数。
291/// 会根据 `db_type` 修正 DSN 的 scheme 后再连接。
292///
293/// # 示例
294///
295/// ```no_run
296/// # #[tokio::main]
297/// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
298/// use sz_orm_core::{Connection, DbType};
299/// let mut conn = sz_orm_sqlx::create_connection_by_type(
300///     DbType::Sqlite, "sqlite::memory:"
301/// ).await?;
302/// conn.execute("SELECT 1").await?;
303/// # Ok(())
304/// # }
305/// ```
306pub 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
320/// 根据 [`AnyBackend`] 修正 DSN 的 scheme
321///
322/// 如果 DSN 已有正确的 scheme 则原样返回;
323/// 否则替换 scheme 部分以匹配目标后端。
324fn 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/// 解析 Oracle DSN 为 (username, password, connect_string)
355///
356/// 格式:`oracle://user:pass@host:port/service`
357///
358/// 返回:username="user", password="pass", connect_string="host:port/service"
359#[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/// 解析 MSSQL DSN 为 ADO.NET 连接字符串
368///
369/// 格式:`mssql://user:pass@host:port/database` 或 `sqlserver://user:pass@host:port/database`
370///
371/// 返回:`Server=host,port;Database=database;User Id=user;Password=pass;`
372#[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/// 通用 DSN 解析:`user:pass@host:port/...` → (user, pass, host:port/...)
390#[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
405/// 后端无关的连接
406pub struct AnyConnection {
407    backend: AnyBackend,
408    inner: Box<dyn Connection>,
409}
410
411impl AnyConnection {
412    /// 获取后端类型
413    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    // ---- 真实 SQLite 集成测试 ----
584
585    #[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        // 创建表并插入数据
594        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        // 查询验证
605        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        // 事务提交
624        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        // 通过 SELECT * 验证行数(避开 sqlx 适配器中 COUNT(*) 类型推断的既有问题)
686        let rows = conn.query("SELECT * FROM cnt").await.unwrap();
687        assert_eq!(rows.len(), 5);
688    }
689
690    // ---- v2.2.0 A-2: Dialect 与 AnyPool 集成测试 ----
691
692    #[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        // 调用各 Dialect trait 方法,确认不 panic(无 todo!/unimplemented!)
750        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    // ---- v6.9.0: from_db_type / create_connection / create_connection_by_type ----
787
788    #[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}