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 }