1use 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 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#[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
51const NO_BACKING: &str = "transaction has no backing connection (created via Transaction::new)";
55
56pub struct Transaction {
57 committed: bool,
58 rolled_back: bool,
59 dialect: Dialect,
62 inner: tokio::sync::Mutex<Option<Box<dyn TransactionInner>>>,
63}
64
65impl Transaction {
66 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 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 async fn execute(&self, sql: &str) -> Result<u64, RdbmsError>;
176 async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
178 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 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 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 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 #[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 tx.commit().await.unwrap();
356 }
357
358 #[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); }
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 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}