1use 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 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#[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 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 async fn execute(&self, sql: &str) -> Result<u64, RdbmsError>;
90 async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
92 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 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 #[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}