reinhardt_core/signals/
transaction.rs1use super::core::SignalName;
28use super::error::SignalError;
29use super::registry::get_signal;
30use super::signal::Signal;
31use serde::{Deserialize, Serialize};
32
33#[derive(Debug, Clone, Serialize, Deserialize)]
45pub struct TransactionContext {
46 pub transaction_id: String,
48 pub savepoint_depth: usize,
50 pub savepoint_name: Option<String>,
52 pub is_nested: bool,
54}
55
56impl TransactionContext {
57 pub fn new(transaction_id: impl Into<String>) -> Self {
69 Self {
70 transaction_id: transaction_id.into(),
71 savepoint_depth: 0,
72 savepoint_name: None,
73 is_nested: false,
74 }
75 }
76
77 pub fn nested(
90 transaction_id: impl Into<String>,
91 depth: usize,
92 savepoint_name: impl Into<String>,
93 ) -> Self {
94 Self {
95 transaction_id: transaction_id.into(),
96 savepoint_depth: depth,
97 savepoint_name: Some(savepoint_name.into()),
98 is_nested: true,
99 }
100 }
101
102 pub fn enter_savepoint(&mut self, name: impl Into<String>) {
115 self.savepoint_depth += 1;
116 self.savepoint_name = Some(name.into());
117 self.is_nested = true;
118 }
119
120 pub fn exit_savepoint(&mut self) {
132 if self.savepoint_depth > 0 {
133 self.savepoint_depth -= 1;
134 }
135 if self.savepoint_depth == 0 {
136 self.savepoint_name = None;
137 self.is_nested = false;
138 }
139 }
140}
141
142pub struct TransactionSignals {
154 context: TransactionContext,
155}
156
157impl TransactionSignals {
158 pub fn new(transaction_id: impl Into<String>) -> Self {
168 Self {
169 context: TransactionContext::new(transaction_id),
170 }
171 }
172
173 pub fn nested(
183 transaction_id: impl Into<String>,
184 depth: usize,
185 savepoint_name: impl Into<String>,
186 ) -> Self {
187 Self {
188 context: TransactionContext::nested(transaction_id, depth, savepoint_name),
189 }
190 }
191
192 pub fn context(&self) -> &TransactionContext {
204 &self.context
205 }
206
207 pub async fn send_begin(&self) -> Result<(), SignalError> {
222 on_begin().send(self.context.clone()).await
223 }
224
225 pub async fn send_commit(&self) -> Result<(), SignalError> {
240 on_commit().send(self.context.clone()).await
241 }
242
243 pub async fn send_rollback(&self) -> Result<(), SignalError> {
258 on_rollback().send(self.context.clone()).await
259 }
260
261 pub async fn enter_savepoint(&mut self, name: impl Into<String>) -> Result<(), SignalError> {
276 self.context.enter_savepoint(name);
277 on_savepoint().send(self.context.clone()).await
278 }
279
280 pub async fn exit_savepoint(&mut self) -> Result<(), SignalError> {
295 self.context.exit_savepoint();
296 on_savepoint_release().send(self.context.clone()).await
297 }
298}
299
300pub fn on_begin() -> Signal<TransactionContext> {
310 get_signal::<TransactionContext>(SignalName::custom("transaction_begin"))
311}
312
313pub fn on_commit() -> Signal<TransactionContext> {
323 get_signal::<TransactionContext>(SignalName::custom("transaction_commit"))
324}
325
326pub fn on_rollback() -> Signal<TransactionContext> {
336 get_signal::<TransactionContext>(SignalName::custom("transaction_rollback"))
337}
338
339pub fn on_savepoint() -> Signal<TransactionContext> {
349 get_signal::<TransactionContext>(SignalName::custom("transaction_savepoint"))
350}
351
352pub fn on_savepoint_release() -> Signal<TransactionContext> {
362 get_signal::<TransactionContext>(SignalName::custom("transaction_savepoint_release"))
363}
364
365#[cfg(test)]
366mod tests {
367 use super::*;
368 use parking_lot::Mutex;
369 use std::sync::Arc;
370 use std::sync::atomic::{AtomicUsize, Ordering};
371
372 #[test]
373 fn test_transaction_context_creation() {
374 let ctx = TransactionContext::new("tx_1");
375 assert_eq!(ctx.transaction_id, "tx_1");
376 assert_eq!(ctx.savepoint_depth, 0);
377 assert_eq!(ctx.savepoint_name, None);
378 assert!(!ctx.is_nested);
379 }
380
381 #[test]
382 fn test_transaction_context_nested() {
383 let ctx = TransactionContext::nested("tx_1", 2, "sp_2");
384 assert_eq!(ctx.transaction_id, "tx_1");
385 assert_eq!(ctx.savepoint_depth, 2);
386 assert_eq!(ctx.savepoint_name, Some("sp_2".to_string()));
387 assert!(ctx.is_nested);
388 }
389
390 #[test]
391 fn test_transaction_context_enter_savepoint() {
392 let mut ctx = TransactionContext::new("tx_1");
393 ctx.enter_savepoint("checkpoint_1");
394 assert_eq!(ctx.savepoint_depth, 1);
395 assert_eq!(ctx.savepoint_name, Some("checkpoint_1".to_string()));
396 assert!(ctx.is_nested);
397 }
398
399 #[test]
400 fn test_transaction_context_exit_savepoint() {
401 let mut ctx = TransactionContext::nested("tx_1", 2, "sp_2");
402 ctx.exit_savepoint();
403 assert_eq!(ctx.savepoint_depth, 1);
404
405 ctx.exit_savepoint();
406 assert_eq!(ctx.savepoint_depth, 0);
407 assert_eq!(ctx.savepoint_name, None);
408 assert!(!ctx.is_nested);
409 }
410
411 #[tokio::test]
412 async fn test_on_commit_signal() {
413 let counter = Arc::new(AtomicUsize::new(0));
414 let counter_clone = Arc::clone(&counter);
415
416 on_commit().connect(move |_ctx| {
417 let counter = Arc::clone(&counter_clone);
418 async move {
419 counter.fetch_add(1, Ordering::SeqCst);
420 Ok(())
421 }
422 });
423
424 let ctx = TransactionContext::new("tx_1");
425 on_commit().send(ctx).await.unwrap();
426
427 assert_eq!(counter.load(Ordering::SeqCst), 1);
428 }
429
430 #[tokio::test]
431 async fn test_on_rollback_signal() {
432 let counter = Arc::new(AtomicUsize::new(0));
433 let counter_clone = Arc::clone(&counter);
434
435 on_rollback().connect(move |_ctx| {
436 let counter = Arc::clone(&counter_clone);
437 async move {
438 counter.fetch_add(1, Ordering::SeqCst);
439 Ok(())
440 }
441 });
442
443 let ctx = TransactionContext::new("tx_1");
444 on_rollback().send(ctx).await.unwrap();
445
446 assert_eq!(counter.load(Ordering::SeqCst), 1);
447 }
448
449 #[tokio::test]
450 #[serial_test::serial]
451 async fn test_transaction_signals_flow() {
452 on_begin().disconnect_all();
454 on_commit().disconnect_all();
455
456 let events = Arc::new(Mutex::new(Vec::new()));
457
458 let e1 = events.clone();
459 on_begin().connect(move |ctx| {
460 let e = e1.clone();
461 async move {
462 e.lock().push(format!("begin:{}", ctx.transaction_id));
463 Ok(())
464 }
465 });
466
467 let e2 = events.clone();
468 on_commit().connect(move |ctx| {
469 let e = e2.clone();
470 async move {
471 e.lock().push(format!("commit:{}", ctx.transaction_id));
472 Ok(())
473 }
474 });
475
476 let signals = TransactionSignals::new("tx_test");
477 signals.send_begin().await.unwrap();
478 signals.send_commit().await.unwrap();
479
480 let event_log = events.lock();
481 assert_eq!(event_log.len(), 2);
482 assert_eq!(event_log[0], "begin:tx_test");
483 assert_eq!(event_log[1], "commit:tx_test");
484
485 on_begin().disconnect_all();
487 on_commit().disconnect_all();
488 }
489
490 #[tokio::test]
491 #[serial_test::serial]
492 async fn test_savepoint_signals() {
493 on_savepoint().disconnect_all();
495 on_savepoint_release().disconnect_all();
496
497 let counter = Arc::new(AtomicUsize::new(0));
498
499 let c1 = counter.clone();
500 on_savepoint().connect(move |_| {
501 let c = c1.clone();
502 async move {
503 c.fetch_add(1, Ordering::SeqCst);
504 Ok(())
505 }
506 });
507
508 let c2 = counter.clone();
509 on_savepoint_release().connect(move |_| {
510 let c = c2.clone();
511 async move {
512 c.fetch_add(10, Ordering::SeqCst);
513 Ok(())
514 }
515 });
516
517 let mut signals = TransactionSignals::new("tx_1");
518 signals.enter_savepoint("sp_1").await.unwrap();
519 signals.exit_savepoint().await.unwrap();
520
521 assert_eq!(counter.load(Ordering::SeqCst), 11); on_savepoint().disconnect_all();
525 on_savepoint_release().disconnect_all();
526 }
527
528 #[tokio::test]
529 #[serial_test::serial]
530 async fn test_nested_transaction_signals() {
531 on_savepoint().disconnect_all();
533
534 let events = Arc::new(Mutex::new(Vec::new()));
535
536 let e = events.clone();
537 on_savepoint().connect(move |ctx| {
538 let e = e.clone();
539 async move {
540 e.lock().push(format!(
541 "savepoint:{}:depth:{}",
542 ctx.savepoint_name.as_deref().unwrap_or(""),
543 ctx.savepoint_depth
544 ));
545 Ok(())
546 }
547 });
548
549 let mut signals = TransactionSignals::new("tx_nested");
550 signals.enter_savepoint("level_1").await.unwrap();
551 signals.enter_savepoint("level_2").await.unwrap();
552
553 let event_log = events.lock();
554 assert_eq!(event_log.len(), 2);
555 assert_eq!(event_log[0], "savepoint:level_1:depth:1");
556 assert_eq!(event_log[1], "savepoint:level_2:depth:2");
557
558 on_savepoint().disconnect_all();
560 }
561}