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
4#[derive(Debug, Clone)]
5pub struct Row {
6    columns: Vec<String>,
7    values: Vec<serde_json::Value>,
8}
9
10impl Row {
11    /// Create a new Row with the given columns and values.
12    pub fn new(columns: Vec<String>, values: Vec<serde_json::Value>) -> Self {
13        debug_assert_eq!(
14            columns.len(),
15            values.len(),
16            "columns and values must have the same length"
17        );
18        Self { columns, values }
19    }
20
21    pub fn get(&self, col: &str) -> Option<&serde_json::Value> {
22        self.columns
23            .iter()
24            .position(|c| c == col)
25            .and_then(|i| self.values.get(i))
26    }
27}
28
29/// Inner transaction trait for cross-backend transaction support.
30#[async_trait]
31pub trait TransactionInner: Send {
32    async fn commit(&mut self) -> Result<(), RdbmsError>;
33    async fn rollback(&mut self) -> Result<(), RdbmsError>;
34}
35
36#[derive(Default)]
37pub struct Transaction {
38    committed: bool,
39    rolled_back: bool,
40    inner: Option<Box<dyn TransactionInner>>,
41}
42
43impl Transaction {
44    pub fn new() -> Self {
45        Self::default()
46    }
47
48    pub fn with_inner(inner: Box<dyn TransactionInner>) -> Self {
49        Self {
50            inner: Some(inner),
51            committed: false,
52            rolled_back: false,
53        }
54    }
55
56    pub async fn commit(mut self) -> Result<(), RdbmsError> {
57        if let Some(ref mut inner) = self.inner {
58            inner.commit().await?;
59        }
60        self.committed = true;
61        Ok(())
62    }
63
64    pub async fn rollback(mut self) -> Result<(), RdbmsError> {
65        if let Some(ref mut inner) = self.inner {
66            inner.rollback().await?;
67        }
68        self.committed = false;
69        self.rolled_back = true;
70        Ok(())
71    }
72}
73
74impl Drop for Transaction {
75    fn drop(&mut self) {
76        // This Drop impl only logs. No SQL is sent here (async work is not
77        // possible in Drop); actual rollback relies on the backing sqlx
78        // Transaction dropping without commit, which rolls back the
79        // underlying DB connection.
80        if !self.committed && !self.rolled_back {
81            tracing::warn!("transaction dropped without commit — rolling back");
82        }
83    }
84}
85
86#[async_trait]
87pub trait RdbmsClient: Send + Sync {
88    /// Execute a raw SQL statement. Prefer `execute_with` for user-supplied values.
89    async fn execute(&self, sql: &str) -> Result<u64, RdbmsError>;
90    /// Query rows with raw SQL. Prefer `query_with` for user-supplied values.
91    async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
92    /// Execute a parameterized SQL statement to prevent injection.
93    /// Backends that cannot bind parameters return an error.
94    async fn execute_with(
95        &self,
96        _sql: &str,
97        _params: &[serde_json::Value],
98    ) -> Result<u64, RdbmsError> {
99        Err(RdbmsError::Database(
100            "parameterized execute not supported by this backend".into(),
101        ))
102    }
103    /// Query with parameterized SQL to prevent injection.
104    /// Backends that cannot bind parameters return an error.
105    async fn query_with(
106        &self,
107        _sql: &str,
108        _params: &[serde_json::Value],
109    ) -> Result<Vec<Row>, RdbmsError> {
110        Err(RdbmsError::Database(
111            "parameterized query not supported by this backend".into(),
112        ))
113    }
114    async fn transaction(&self) -> Result<Transaction, RdbmsError>;
115}
116
117#[derive(Debug, thiserror::Error)]
118pub enum RdbmsError {
119    #[error("database error: {0}")]
120    Database(String),
121    #[error("connection error: {0}")]
122    Connection(String),
123    #[error("configuration error: {0}")]
124    Config(String),
125}
126
127#[cfg(test)]
128mod tests {
129    use super::*;
130    use std::sync::Arc;
131    use std::sync::atomic::{AtomicUsize, Ordering};
132
133    #[test]
134    fn row_get_returns_value_by_column() {
135        let row = Row::new(
136            vec!["id".into(), "name".into()],
137            vec![serde_json::json!(1), serde_json::json!("alice")],
138        );
139        assert_eq!(row.get("name"), Some(&serde_json::json!("alice")));
140        assert_eq!(row.get("missing"), None);
141    }
142
143    #[test]
144    fn row_get_uses_first_matching_column() {
145        let row = Row::new(
146            vec!["a".into(), "a".into()],
147            vec![serde_json::json!(1), serde_json::json!(2)],
148        );
149        assert_eq!(row.get("a"), Some(&serde_json::json!(1)));
150    }
151
152    #[derive(Clone, Default)]
153    struct Tracked {
154        commits: Arc<AtomicUsize>,
155        rollbacks: Arc<AtomicUsize>,
156    }
157
158    struct TrackingInner {
159        track: Tracked,
160    }
161
162    #[async_trait]
163    impl TransactionInner for TrackingInner {
164        async fn commit(&mut self) -> Result<(), RdbmsError> {
165            self.track.commits.fetch_add(1, Ordering::SeqCst);
166            Ok(())
167        }
168        async fn rollback(&mut self) -> Result<(), RdbmsError> {
169            self.track.rollbacks.fetch_add(1, Ordering::SeqCst);
170            Ok(())
171        }
172    }
173
174    #[tokio::test]
175    async fn commit_delegates_to_inner() {
176        let track = Tracked::default();
177        let tx = Transaction::with_inner(Box::new(TrackingInner {
178            track: track.clone(),
179        }));
180        tx.commit().await.unwrap();
181        assert_eq!(track.commits.load(Ordering::SeqCst), 1);
182        assert_eq!(track.rollbacks.load(Ordering::SeqCst), 0);
183    }
184
185    #[tokio::test]
186    async fn rollback_delegates_to_inner() {
187        let track = Tracked::default();
188        let tx = Transaction::with_inner(Box::new(TrackingInner {
189            track: track.clone(),
190        }));
191        tx.rollback().await.unwrap();
192        assert_eq!(track.rollbacks.load(Ordering::SeqCst), 1);
193        assert_eq!(track.commits.load(Ordering::SeqCst), 0);
194    }
195
196    #[tokio::test]
197    async fn commit_without_inner_succeeds() {
198        let tx = Transaction::new();
199        tx.commit().await.unwrap();
200    }
201
202    /// 只统计 WARN 事件的最小 Subscriber,用于验证 Drop guard 的告警行为。
203    #[derive(Clone)]
204    struct WarnCounter(Arc<AtomicUsize>);
205
206    impl tracing::Subscriber for WarnCounter {
207        fn enabled(&self, _: &tracing::Metadata<'_>) -> bool {
208            true
209        }
210        fn new_span(&self, _: &tracing::span::Attributes<'_>) -> tracing::span::Id {
211            tracing::span::Id::from_u64(1)
212        }
213        fn record(&self, _: &tracing::span::Id, _: &tracing::span::Record<'_>) {}
214        fn record_follows_from(&self, _: &tracing::span::Id, _: &tracing::span::Id) {}
215        fn event(&self, event: &tracing::Event<'_>) {
216            if *event.metadata().level() == tracing::Level::WARN {
217                self.0.fetch_add(1, Ordering::SeqCst);
218            }
219        }
220        fn enter(&self, _: &tracing::span::Id) {}
221        fn exit(&self, _: &tracing::span::Id) {}
222    }
223
224    fn with_warn_counter(counts: Arc<AtomicUsize>, f: impl FnOnce()) {
225        tracing::subscriber::with_default(WarnCounter(counts), f);
226    }
227
228    #[test]
229    fn drop_after_explicit_rollback_does_not_warn() {
230        let warns = Arc::new(AtomicUsize::new(0));
231        let track = Tracked::default();
232        let tx = Transaction::with_inner(Box::new(TrackingInner {
233            track: track.clone(),
234        }));
235        with_warn_counter(Arc::clone(&warns), || {
236            tokio::runtime::Builder::new_current_thread()
237                .build()
238                .unwrap()
239                .block_on(tx.rollback())
240                .unwrap();
241        });
242        assert_eq!(track.rollbacks.load(Ordering::SeqCst), 1);
243        assert_eq!(track.commits.load(Ordering::SeqCst), 0);
244        assert_eq!(
245            warns.load(Ordering::SeqCst),
246            0,
247            "rollback 后 Drop 不得再告警"
248        );
249    }
250
251    #[test]
252    fn drop_without_commit_warns_once() {
253        let warns = Arc::new(AtomicUsize::new(0));
254        let tx = Transaction::with_inner(Box::new(TrackingInner {
255            track: Tracked::default(),
256        }));
257        with_warn_counter(Arc::clone(&warns), || drop(tx));
258        assert_eq!(
259            warns.load(Ordering::SeqCst),
260            1,
261            "未提交即 Drop 必须告警一次"
262        );
263    }
264
265    struct RawOnlyClient;
266
267    #[async_trait]
268    impl RdbmsClient for RawOnlyClient {
269        async fn execute(&self, _sql: &str) -> Result<u64, RdbmsError> {
270            Ok(0)
271        }
272        async fn query(&self, _sql: &str) -> Result<Vec<Row>, RdbmsError> {
273            Ok(vec![])
274        }
275        async fn transaction(&self) -> Result<Transaction, RdbmsError> {
276            Ok(Transaction::new())
277        }
278    }
279
280    #[tokio::test]
281    async fn parameterized_ops_default_to_not_supported_error() {
282        let client = RawOnlyClient;
283        let err = client.execute_with("SELECT 1", &[]).await.unwrap_err();
284        assert!(
285            err.to_string()
286                .contains("parameterized execute not supported"),
287            "got: {err}"
288        );
289        let err = client.query_with("SELECT 1", &[]).await.unwrap_err();
290        assert!(
291            err.to_string()
292                .contains("parameterized query not supported"),
293            "got: {err}"
294        );
295    }
296}