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::{Connection, ConnectionFactory, DbError, QueryRows};
35
36use crate::any::{
37    MySqlPoolHandle, PgPoolHandle, SqlitePoolHandle, SqlxMySqlConnectionFactory,
38    SqlxPgConnectionFactory, SqlxSqliteConnectionFactory,
39};
40
41/// 数据库后端类型
42#[derive(Debug, Clone, Copy, PartialEq, Eq)]
43pub enum AnyBackend {
44    /// MySQL / MariaDB
45    MySql,
46    /// PostgreSQL
47    Postgres,
48    /// SQLite
49    Sqlite,
50}
51
52impl AnyBackend {
53    /// 从 DSN 自动识别后端类型
54    ///
55    /// # 支持的 scheme
56    ///
57    /// - `mysql://` / `mariadb://` → MySQL
58    /// - `postgres://` / `postgresql://` → Postgres
59    /// - `sqlite://` / `sqlite:` → Sqlite
60    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    /// 后端名称
76    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
85/// 后端无关的连接工厂
86pub struct AnyPool {
87    backend: AnyBackend,
88    factory: Arc<dyn ConnectionFactory>,
89}
90
91impl AnyPool {
92    /// 连接数据库,根据 DSN 自动识别后端
93    ///
94    /// # 错误
95    ///
96    /// - DSN scheme 不识别 → [`DbError::ConnectionRefused`]
97    /// - 连接失败 → [`DbError::ConnectionError`]
98    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    /// 从已有的连接工厂构造
118    pub fn from_factory(backend: AnyBackend, factory: Arc<dyn ConnectionFactory>) -> Self {
119        Self { backend, factory }
120    }
121
122    /// 获取后端类型
123    pub fn backend(&self) -> AnyBackend {
124        self.backend
125    }
126
127    /// 创建一个新连接
128    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
137/// 后端无关的连接
138pub struct AnyConnection {
139    backend: AnyBackend,
140    inner: Box<dyn Connection>,
141}
142
143impl AnyConnection {
144    /// 获取后端类型
145    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    // ---- 真实 SQLite 集成测试 ----
259
260    #[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        // 创建表并插入数据
269        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        // 查询验证
280        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        // 事务提交
299        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        // 通过 SELECT * 验证行数(避开 sqlx 适配器中 COUNT(*) 类型推断的既有问题)
361        let rows = conn.query("SELECT * FROM cnt").await.unwrap();
362        assert_eq!(rows.len(), 5);
363    }
364}