lib

Core libraries for Radroots
git clone https://radroots.dev/git/lib.git
Log | Files | Refs | README

transaction_control.rs (7314B)


      1 //! Private deny-by-default SQLite transaction-control fencing.
      2 
      3 #[cfg(any(target_os = "linux", target_os = "macos"))]
      4 use std::sync::{
      5     Arc,
      6     atomic::{AtomicBool, Ordering},
      7 };
      8 
      9 #[cfg(any(target_os = "linux", target_os = "macos"))]
     10 use sqlx::SqliteConnection;
     11 
     12 #[cfg(any(target_os = "linux", target_os = "macos"))]
     13 pub(crate) struct TransactionControlGate {
     14     allow_commit: Arc<AtomicBool>,
     15     allow_runner_rollback: Arc<AtomicBool>,
     16     rejected_commit: Arc<AtomicBool>,
     17     rollback_observed: Arc<AtomicBool>,
     18 }
     19 
     20 #[cfg(any(target_os = "linux", target_os = "macos"))]
     21 impl TransactionControlGate {
     22     pub(crate) async fn install(connection: &mut SqliteConnection) -> Result<Self, sqlx::Error> {
     23         let allow_commit = Arc::new(AtomicBool::new(false));
     24         let allow_runner_rollback = Arc::new(AtomicBool::new(false));
     25         let rejected_commit = Arc::new(AtomicBool::new(false));
     26         let rollback_observed = Arc::new(AtomicBool::new(false));
     27         let hook_permission = Arc::clone(&allow_commit);
     28         let rejected_commit_epoch = Arc::clone(&rejected_commit);
     29         let rollback_permission = Arc::clone(&allow_runner_rollback);
     30         let rollback_epoch = Arc::clone(&rollback_observed);
     31         let mut handle = connection.lock_handle().await?;
     32         handle.set_commit_hook(move || {
     33             let permitted = hook_permission.load(Ordering::Acquire);
     34             if !permitted {
     35                 rejected_commit_epoch.store(true, Ordering::Release);
     36             }
     37             permitted
     38         });
     39         handle.set_rollback_hook(move || {
     40             if !rollback_permission.load(Ordering::Acquire) {
     41                 rollback_epoch.store(true, Ordering::Release);
     42             }
     43         });
     44         drop(handle);
     45         Ok(Self {
     46             allow_commit,
     47             allow_runner_rollback,
     48             rejected_commit,
     49             rollback_observed,
     50         })
     51     }
     52 
     53     pub(crate) fn permit_outer_commit(&self) -> TransactionCommitPermit {
     54         self.allow_commit.store(true, Ordering::Release);
     55         TransactionCommitPermit {
     56             allow_commit: Arc::clone(&self.allow_commit),
     57         }
     58     }
     59 
     60     pub(crate) fn permit_runner_rollback(&self) -> TransactionRollbackPermit {
     61         self.allow_runner_rollback.store(true, Ordering::Release);
     62         TransactionRollbackPermit {
     63             allow_runner_rollback: Arc::clone(&self.allow_runner_rollback),
     64         }
     65     }
     66 
     67     pub(crate) fn control_violation_observed(&self) -> bool {
     68         self.rejected_commit.load(Ordering::Acquire)
     69             || self.rollback_observed.load(Ordering::Acquire)
     70     }
     71 
     72     pub(crate) fn rejected_commit_rolled_back(&self) -> bool {
     73         self.rejected_commit.load(Ordering::Acquire)
     74             && self.rollback_observed.load(Ordering::Acquire)
     75     }
     76 
     77     pub(crate) async fn remove(self, connection: &mut SqliteConnection) -> Result<(), sqlx::Error> {
     78         self.allow_commit.store(false, Ordering::Release);
     79         self.allow_runner_rollback.store(false, Ordering::Release);
     80         let mut handle = connection.lock_handle().await?;
     81         handle.remove_commit_hook();
     82         handle.remove_rollback_hook();
     83         Ok(())
     84     }
     85 }
     86 
     87 #[cfg(any(target_os = "linux", target_os = "macos"))]
     88 pub(crate) struct TransactionCommitPermit {
     89     allow_commit: Arc<AtomicBool>,
     90 }
     91 
     92 #[cfg(any(target_os = "linux", target_os = "macos"))]
     93 impl Drop for TransactionCommitPermit {
     94     fn drop(&mut self) {
     95         self.allow_commit.store(false, Ordering::Release);
     96     }
     97 }
     98 
     99 #[cfg(any(target_os = "linux", target_os = "macos"))]
    100 pub(crate) struct TransactionRollbackPermit {
    101     allow_runner_rollback: Arc<AtomicBool>,
    102 }
    103 
    104 #[cfg(any(target_os = "linux", target_os = "macos"))]
    105 impl Drop for TransactionRollbackPermit {
    106     fn drop(&mut self) {
    107         self.allow_runner_rollback.store(false, Ordering::Release);
    108     }
    109 }
    110 
    111 #[cfg(all(test, any(target_os = "linux", target_os = "macos")))]
    112 mod tests {
    113     use sqlx::{Connection, SqliteConnection};
    114 
    115     use super::TransactionControlGate;
    116 
    117     async fn memory_connection() -> SqliteConnection {
    118         let mut connection = SqliteConnection::connect("sqlite::memory:")
    119             .await
    120             .expect("memory SQLite connection");
    121         sqlx::query("CREATE TABLE gate_probe (value INTEGER NOT NULL)")
    122             .execute(&mut connection)
    123             .await
    124             .expect("gate probe table");
    125         connection
    126     }
    127 
    128     #[tokio::test(flavor = "current_thread")]
    129     async fn denied_commit_and_unpermitted_rollback_are_observed_exactly() {
    130         let mut connection = memory_connection().await;
    131         let gate = TransactionControlGate::install(&mut connection)
    132             .await
    133             .expect("transaction gate");
    134         let mut transaction = connection.begin().await.expect("transaction");
    135         sqlx::query("INSERT INTO gate_probe (value) VALUES (1)")
    136             .execute(&mut *transaction)
    137             .await
    138             .expect("mutate denied transaction");
    139         transaction
    140             .commit()
    141             .await
    142             .expect_err("commit must be denied");
    143         assert!(gate.control_violation_observed());
    144         assert!(gate.rejected_commit_rolled_back());
    145         gate.remove(&mut connection).await.expect("remove gate");
    146 
    147         let mut connection = memory_connection().await;
    148         let gate = TransactionControlGate::install(&mut connection)
    149             .await
    150             .expect("transaction gate");
    151         let mut transaction = connection.begin().await.expect("transaction");
    152         sqlx::query("INSERT INTO gate_probe (value) VALUES (2)")
    153             .execute(&mut *transaction)
    154             .await
    155             .expect("mutate rolled back transaction");
    156         transaction.rollback().await.expect("SQLite rollback");
    157         assert!(gate.control_violation_observed());
    158         assert!(!gate.rejected_commit_rolled_back());
    159         gate.remove(&mut connection).await.expect("remove gate");
    160     }
    161 
    162     #[tokio::test(flavor = "current_thread")]
    163     async fn runner_permits_are_scoped_and_do_not_record_violations() {
    164         let mut connection = memory_connection().await;
    165         let gate = TransactionControlGate::install(&mut connection)
    166             .await
    167             .expect("transaction gate");
    168 
    169         let mut transaction = connection.begin().await.expect("transaction");
    170         sqlx::query("INSERT INTO gate_probe (value) VALUES (3)")
    171             .execute(&mut *transaction)
    172             .await
    173             .expect("mutate committed transaction");
    174         let permit = gate.permit_outer_commit();
    175         transaction.commit().await.expect("permitted commit");
    176         drop(permit);
    177         assert!(!gate.control_violation_observed());
    178         assert!(!gate.rejected_commit_rolled_back());
    179 
    180         let mut transaction = connection.begin().await.expect("transaction");
    181         sqlx::query("INSERT INTO gate_probe (value) VALUES (4)")
    182             .execute(&mut *transaction)
    183             .await
    184             .expect("mutate runner rollback transaction");
    185         let permit = gate.permit_runner_rollback();
    186         transaction.rollback().await.expect("permitted rollback");
    187         drop(permit);
    188         assert!(!gate.control_violation_observed());
    189         assert!(!gate.rejected_commit_rolled_back());
    190         gate.remove(&mut connection).await.expect("remove gate");
    191     }
    192 }