Skip to main content

ecat_data/
rdbms.rs

1// Copyright (c) 2026 erik <erik@erik.xyz> — https://erik.xyz
2use async_trait::async_trait;
3
4use crate::dialect::Dialect;
5
6#[derive(Debug, Clone)]
7pub struct Row {
8    columns: Vec<String>,
9    values: Vec<serde_json::Value>,
10}
11
12impl Row {
13    /// Create a new Row with the given columns and values.
14    pub fn new(columns: Vec<String>, values: Vec<serde_json::Value>) -> Self {
15        debug_assert_eq!(
16            columns.len(),
17            values.len(),
18            "columns and values must have the same length"
19        );
20        Self { columns, values }
21    }
22
23    pub fn get(&self, col: &str) -> Option<&serde_json::Value> {
24        self.columns
25            .iter()
26            .position(|c| c == col)
27            .and_then(|i| self.values.get(i))
28    }
29}
30
31/// 事务内部实现。后端(sqlx / tiberius)实现它,`Transaction` 转发调用。
32#[async_trait]
33pub trait TransactionInner: Send {
34    async fn execute(&mut self, sql: &str) -> Result<u64, RdbmsError>;
35    async fn query(&mut self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
36    async fn execute_with(
37        &mut self,
38        sql: &str,
39        params: &[serde_json::Value],
40    ) -> Result<u64, RdbmsError>;
41    async fn query_with(
42        &mut self,
43        sql: &str,
44        params: &[serde_json::Value],
45    ) -> Result<Vec<Row>, RdbmsError>;
46    fn dialect(&self) -> Dialect;
47    async fn commit(&mut self) -> Result<(), RdbmsError>;
48    async fn rollback(&mut self) -> Result<(), RdbmsError>;
49}
50
51/// 无 backing 连接的事务([`Transaction::new`])执行 SQL 时的错误文案。
52/// 这类事务只能作为空占位,执行任何语句都是编程错误 —— 必须报错而非
53/// 静默返回 0 行影响,否则写操作会无声丢失。
54const NO_BACKING: &str = "transaction has no backing connection (created via Transaction::new)";
55
56pub struct Transaction {
57    committed: bool,
58    rolled_back: bool,
59    /// 在 `with_inner` 时从 inner 拷贝,避免 `dialect(&self)` 这个同步方法
60    /// 需要等待异步锁。
61    dialect: Dialect,
62    inner: tokio::sync::Mutex<Option<Box<dyn TransactionInner>>>,
63}
64
65impl Transaction {
66    /// 创建一个无 backing 连接的空事务。
67    ///
68    /// 只能作为占位符用于「不执行任何语句」的场景;在其中执行 SQL 会返回错误。
69    /// 需要真正执行语句时用 [`Transaction::with_inner`] 或后端的 `transaction()`。
70    pub fn new() -> Self {
71        Self {
72            committed: false,
73            rolled_back: false,
74            dialect: Dialect::Standard,
75            inner: tokio::sync::Mutex::new(None),
76        }
77    }
78
79    pub fn with_inner(inner: Box<dyn TransactionInner>) -> Self {
80        let dialect = inner.dialect();
81        Self {
82            committed: false,
83            rolled_back: false,
84            dialect,
85            inner: tokio::sync::Mutex::new(Some(inner)),
86        }
87    }
88
89    pub async fn commit(mut self) -> Result<(), RdbmsError> {
90        if let Some(inner) = self.inner.get_mut().as_mut() {
91            inner.commit().await?;
92        }
93        self.committed = true;
94        Ok(())
95    }
96
97    pub async fn rollback(mut self) -> Result<(), RdbmsError> {
98        if let Some(inner) = self.inner.get_mut().as_mut() {
99            inner.rollback().await?;
100        }
101        self.committed = false;
102        self.rolled_back = true;
103        Ok(())
104    }
105}
106
107impl Default for Transaction {
108    fn default() -> Self {
109        Self::new()
110    }
111}
112
113#[async_trait]
114impl SqlExecutor for Transaction {
115    async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
116        let mut guard = self.inner.lock().await;
117        match guard.as_mut() {
118            Some(inner) => inner.execute(sql).await,
119            None => Err(RdbmsError::Database(NO_BACKING.into())),
120        }
121    }
122
123    async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
124        let mut guard = self.inner.lock().await;
125        match guard.as_mut() {
126            Some(inner) => inner.query(sql).await,
127            None => Err(RdbmsError::Database(NO_BACKING.into())),
128        }
129    }
130
131    async fn execute_with(
132        &self,
133        sql: &str,
134        params: &[serde_json::Value],
135    ) -> Result<u64, RdbmsError> {
136        let mut guard = self.inner.lock().await;
137        match guard.as_mut() {
138            Some(inner) => inner.execute_with(sql, params).await,
139            None => Err(RdbmsError::Database(NO_BACKING.into())),
140        }
141    }
142
143    async fn query_with(
144        &self,
145        sql: &str,
146        params: &[serde_json::Value],
147    ) -> Result<Vec<Row>, RdbmsError> {
148        let mut guard = self.inner.lock().await;
149        match guard.as_mut() {
150            Some(inner) => inner.query_with(sql, params).await,
151            None => Err(RdbmsError::Database(NO_BACKING.into())),
152        }
153    }
154
155    fn dialect(&self) -> Dialect {
156        self.dialect
157    }
158}
159
160impl Drop for Transaction {
161    fn drop(&mut self) {
162        // 这里只记日志与计数:Drop 里无法执行异步回滚,实际回滚依赖
163        // 底层 sqlx / tiberius 事务在未提交时 Drop 自动回滚。
164        if !self.committed && !self.rolled_back {
165            crate::timeout::TRANSACTIONS_LEAKED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
166            tracing::warn!("transaction dropped without commit — rolling back");
167        }
168    }
169}
170
171#[async_trait]
172pub trait SqlExecutor: Send + Sync {
173    /// 执行一条 SQL 语句,返回受影响行数。
174    /// 用户提供的值请走 [`SqlExecutor::execute_with`]。
175    async fn execute(&self, sql: &str) -> Result<u64, RdbmsError>;
176    /// 查询多行。用户提供的值请走 [`SqlExecutor::query_with`]。
177    async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
178    /// 参数化执行,防注入。无法绑定参数的后端返回错误。
179    async fn execute_with(
180        &self,
181        _sql: &str,
182        _params: &[serde_json::Value],
183    ) -> Result<u64, RdbmsError> {
184        Err(RdbmsError::Database(
185            "parameterized execute not supported by this backend".into(),
186        ))
187    }
188    /// 参数化查询,防注入。无法绑定参数的后端返回错误。
189    async fn query_with(
190        &self,
191        _sql: &str,
192        _params: &[serde_json::Value],
193    ) -> Result<Vec<Row>, RdbmsError> {
194        Err(RdbmsError::Database(
195            "parameterized query not supported by this backend".into(),
196        ))
197    }
198    /// 写路径且需要返回结果(`INSERT ... RETURNING` / `OUTPUT INSERTED`)。
199    /// 默认委托给 [`SqlExecutor::query_with`];只有读写分离路由需要覆写,
200    /// 否则写语句会被路由到从库。
201    async fn query_write(
202        &self,
203        sql: &str,
204        params: &[serde_json::Value],
205    ) -> Result<Vec<Row>, RdbmsError> {
206        self.query_with(sql, params).await
207    }
208    /// 本执行器背后的数据库方言。
209    fn dialect(&self) -> Dialect;
210}
211
212#[async_trait]
213pub trait RdbmsClient: SqlExecutor {
214    async fn transaction(&self) -> Result<Transaction, RdbmsError>;
215}
216
217#[derive(Debug, thiserror::Error)]
218pub enum RdbmsError {
219    #[error("database error: {0}")]
220    Database(String),
221    #[error("connection error: {0}")]
222    Connection(String),
223    #[error("configuration error: {0}")]
224    Config(String),
225    #[error("timeout: {0}")]
226    Timeout(String),
227}
228
229#[cfg(test)]
230mod tests {
231    use super::*;
232    use std::sync::Arc;
233    use std::sync::atomic::{AtomicUsize, Ordering};
234
235    #[test]
236    fn row_get_returns_value_by_column() {
237        let row = Row::new(
238            vec!["id".into(), "name".into()],
239            vec![serde_json::json!(1), serde_json::json!("alice")],
240        );
241        assert_eq!(row.get("name"), Some(&serde_json::json!("alice")));
242        assert_eq!(row.get("missing"), None);
243    }
244
245    #[test]
246    fn row_get_uses_first_matching_column() {
247        let row = Row::new(
248            vec!["a".into(), "a".into()],
249            vec![serde_json::json!(1), serde_json::json!(2)],
250        );
251        assert_eq!(row.get("a"), Some(&serde_json::json!(1)));
252    }
253
254    #[derive(Clone, Default)]
255    struct Tracked {
256        commits: Arc<AtomicUsize>,
257        rollbacks: Arc<AtomicUsize>,
258        executes: Arc<AtomicUsize>,
259    }
260
261    struct TrackingInner {
262        track: Tracked,
263    }
264
265    #[async_trait]
266    impl TransactionInner for TrackingInner {
267        async fn execute(&mut self, _sql: &str) -> Result<u64, RdbmsError> {
268            self.track.executes.fetch_add(1, Ordering::SeqCst);
269            Ok(7)
270        }
271        async fn query(&mut self, _sql: &str) -> Result<Vec<Row>, RdbmsError> {
272            Ok(vec![Row::new(vec!["n".into()], vec![serde_json::json!(1)])])
273        }
274        async fn execute_with(
275            &mut self,
276            _sql: &str,
277            _p: &[serde_json::Value],
278        ) -> Result<u64, RdbmsError> {
279            Ok(0)
280        }
281        async fn query_with(
282            &mut self,
283            _sql: &str,
284            _p: &[serde_json::Value],
285        ) -> Result<Vec<Row>, RdbmsError> {
286            Ok(vec![])
287        }
288        fn dialect(&self) -> Dialect {
289            Dialect::Sqlite
290        }
291        async fn commit(&mut self) -> Result<(), RdbmsError> {
292            self.track.commits.fetch_add(1, Ordering::SeqCst);
293            Ok(())
294        }
295        async fn rollback(&mut self) -> Result<(), RdbmsError> {
296            self.track.rollbacks.fetch_add(1, Ordering::SeqCst);
297            Ok(())
298        }
299    }
300
301    #[tokio::test]
302    async fn commit_delegates_to_inner() {
303        let track = Tracked::default();
304        let tx = Transaction::with_inner(Box::new(TrackingInner {
305            track: track.clone(),
306        }));
307        tx.commit().await.unwrap();
308        assert_eq!(track.commits.load(Ordering::SeqCst), 1);
309        assert_eq!(track.rollbacks.load(Ordering::SeqCst), 0);
310    }
311
312    #[tokio::test]
313    async fn rollback_delegates_to_inner() {
314        let track = Tracked::default();
315        let tx = Transaction::with_inner(Box::new(TrackingInner {
316            track: track.clone(),
317        }));
318        tx.rollback().await.unwrap();
319        assert_eq!(track.rollbacks.load(Ordering::SeqCst), 1);
320        assert_eq!(track.commits.load(Ordering::SeqCst), 0);
321    }
322
323    #[tokio::test]
324    async fn commit_without_inner_succeeds() {
325        let tx = Transaction::new();
326        tx.commit().await.unwrap();
327    }
328
329    #[tokio::test]
330    async fn transaction_executes_within_scope() {
331        let track = Tracked::default();
332        let tx = Transaction::with_inner(Box::new(TrackingInner {
333            track: track.clone(),
334        }));
335        assert_eq!(tx.execute("UPDATE t SET x = 1").await.unwrap(), 7);
336        assert_eq!(track.executes.load(Ordering::SeqCst), 1);
337    }
338
339    #[tokio::test]
340    async fn transaction_reports_inner_dialect() {
341        let tx = Transaction::with_inner(Box::new(TrackingInner {
342            track: Tracked::default(),
343        }));
344        assert_eq!(tx.dialect(), Dialect::Sqlite);
345    }
346
347    /// 空事务执行 SQL 必须报错,而不是静默返回 0 行影响 ——
348    /// 后者会让写操作无声丢失(审查发现)。
349    #[tokio::test]
350    async fn empty_transaction_rejects_execution() {
351        let tx = Transaction::new();
352        assert!(tx.execute("SELECT 1").await.is_err());
353        assert!(tx.query("SELECT 1").await.is_err());
354        // 没有东西要提交,不应视为错误
355        tx.commit().await.unwrap();
356    }
357
358    /// 只统计 WARN 事件的最小 Subscriber,用于验证 Drop guard 的告警行为。
359    #[derive(Clone)]
360    struct WarnCounter(Arc<AtomicUsize>);
361
362    impl tracing::Subscriber for WarnCounter {
363        fn enabled(&self, _: &tracing::Metadata<'_>) -> bool {
364            true
365        }
366        fn new_span(&self, _: &tracing::span::Attributes<'_>) -> tracing::span::Id {
367            tracing::span::Id::from_u64(1)
368        }
369        fn record(&self, _: &tracing::span::Id, _: &tracing::span::Record<'_>) {}
370        fn record_follows_from(&self, _: &tracing::span::Id, _: &tracing::span::Id) {}
371        fn event(&self, event: &tracing::Event<'_>) {
372            if *event.metadata().level() == tracing::Level::WARN {
373                self.0.fetch_add(1, Ordering::SeqCst);
374            }
375        }
376        fn enter(&self, _: &tracing::span::Id) {}
377        fn exit(&self, _: &tracing::span::Id) {}
378    }
379
380    fn with_warn_counter(counts: Arc<AtomicUsize>, f: impl FnOnce()) {
381        tracing::subscriber::with_default(WarnCounter(counts), f);
382    }
383
384    #[test]
385    fn drop_after_explicit_rollback_does_not_warn() {
386        let warns = Arc::new(AtomicUsize::new(0));
387        let track = Tracked::default();
388        let tx = Transaction::with_inner(Box::new(TrackingInner {
389            track: track.clone(),
390        }));
391        with_warn_counter(Arc::clone(&warns), || {
392            tokio::runtime::Builder::new_current_thread()
393                .build()
394                .unwrap()
395                .block_on(tx.rollback())
396                .unwrap();
397        });
398        assert_eq!(track.rollbacks.load(Ordering::SeqCst), 1);
399        assert_eq!(track.commits.load(Ordering::SeqCst), 0);
400        assert_eq!(
401            warns.load(Ordering::SeqCst),
402            0,
403            "rollback 后 Drop 不得再告警"
404        );
405    }
406
407    #[test]
408    fn dropped_uncommitted_transaction_still_warns() {
409        let warns = Arc::new(AtomicUsize::new(0));
410        let tx = Transaction::with_inner(Box::new(TrackingInner {
411            track: Tracked::default(),
412        }));
413        with_warn_counter(Arc::clone(&warns), || drop(tx));
414        assert_eq!(warns.load(Ordering::SeqCst), 1);
415    }
416
417    #[test]
418    fn dropped_uncommitted_transaction_counts_as_leak() {
419        use crate::timeout::TRANSACTIONS_LEAKED;
420        let before = TRANSACTIONS_LEAKED.load(Ordering::SeqCst);
421        drop(Transaction::new());
422        assert!(TRANSACTIONS_LEAKED.load(Ordering::SeqCst) > before); // 并行测试也会递增
423    }
424
425    struct RawOnlyClient;
426
427    #[async_trait]
428    impl SqlExecutor for RawOnlyClient {
429        async fn execute(&self, _sql: &str) -> Result<u64, RdbmsError> {
430            Ok(0)
431        }
432        async fn query(&self, _sql: &str) -> Result<Vec<Row>, RdbmsError> {
433            Ok(vec![])
434        }
435        fn dialect(&self) -> Dialect {
436            Dialect::Standard
437        }
438    }
439
440    #[tokio::test]
441    async fn parameterized_ops_default_to_not_supported_error() {
442        let client = RawOnlyClient;
443        let err = client.execute_with("SELECT 1", &[]).await.unwrap_err();
444        assert!(
445            err.to_string()
446                .contains("parameterized execute not supported"),
447            "got: {err}"
448        );
449        let err = client.query_with("SELECT 1", &[]).await.unwrap_err();
450        assert!(
451            err.to_string()
452                .contains("parameterized query not supported"),
453            "got: {err}"
454        );
455    }
456
457    /// 默认的 `query_write` 必须委托给 `query_with`:这样只有读写分离路由
458    /// 需要覆写它,其余后端(含第三方实现)零改动即可支持写路径。
459    struct CountingClient {
460        query_with_calls: Arc<AtomicUsize>,
461    }
462
463    #[async_trait]
464    impl SqlExecutor for CountingClient {
465        async fn execute(&self, _sql: &str) -> Result<u64, RdbmsError> {
466            Ok(0)
467        }
468        async fn query(&self, _sql: &str) -> Result<Vec<Row>, RdbmsError> {
469            Ok(vec![])
470        }
471        async fn query_with(
472            &self,
473            _sql: &str,
474            _params: &[serde_json::Value],
475        ) -> Result<Vec<Row>, RdbmsError> {
476            self.query_with_calls.fetch_add(1, Ordering::SeqCst);
477            Ok(vec![])
478        }
479        fn dialect(&self) -> Dialect {
480            Dialect::Standard
481        }
482    }
483
484    #[tokio::test]
485    async fn query_write_defaults_to_query_with() {
486        let client = CountingClient {
487            query_with_calls: Arc::new(AtomicUsize::new(0)),
488        };
489        client.query_write("SELECT 1", &[]).await.unwrap();
490        assert_eq!(client.query_with_calls.load(Ordering::SeqCst), 1);
491    }
492
493    #[test]
494    fn timeout_error_renders_message() {
495        let err = RdbmsError::Timeout("query exceeded 30s".into());
496        assert!(err.to_string().contains("timeout"));
497        assert!(err.to_string().contains("30s"));
498    }
499}