From 2addc1133dc7bd680d2461caeab4d0b305aac95f Mon Sep 17 00:00:00 2001 From: Leo dev Date: Sat, 31 Jan 2026 12:57:56 +0100 Subject: [PATCH] Thread safe --- src/lib.rs | 103 ++++++++++++++++++++++++++++++++++------------------- 1 file changed, 66 insertions(+), 37 deletions(-) diff --git a/src/lib.rs b/src/lib.rs index 6f6ee67..f6fdbf4 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -3,11 +3,15 @@ //! When you call [`access`](AccessCell::access) with a closure that itself calls `access` //! on the same cell, those nested calls are queued and run after the current closure //! finishes, avoiding deadlock. +//! +//! This type is thread-safe: it implements [`Send`] and [`Sync`] when `T: Send`. use std::{ - cell::{Cell, UnsafeCell}, collections::VecDeque, - sync::Mutex, + sync::{ + atomic::{AtomicBool, Ordering}, + Mutex, RwLock, + }, }; /// A cell holding a value of type `T` with re-entrant mutable access. @@ -15,72 +19,72 @@ use std::{ /// The first caller of [`access`](AccessCell::access) runs immediately; any further /// calls to `access` from within that closure (or its queue) are enqueued and /// executed in order after the current work completes. +/// +/// Thread-safe: safe to share across threads and use from multiple threads. pub struct AccessCell { - value: UnsafeCell, - running: Cell, - queue: Mutex>>, + value: RwLock, + running: AtomicBool, + queue: Mutex>>, } +// Safe: only one thread holds the write guard at a time (enforced by `running` + queue). +unsafe impl Send for AccessCell {} +unsafe impl Sync for AccessCell {} + impl AccessCell { /// Creates a new `AccessCell` wrapping `value`. pub fn new(value: T) -> Self { Self { - value: UnsafeCell::new(value), - running: Cell::new(false), + value: RwLock::new(value), + running: AtomicBool::new(false), queue: Mutex::new(VecDeque::new()), } } /// Runs `f` with exclusive mutable access to the inner value. /// - /// If called while another `access` closure is running (re-entrant call), - /// `f` is queued and run after the current closure and any other queued - /// closures finish. - pub fn access(&self, f: impl FnOnce(&mut T) + 'static) { - // already inside → enqueue - if self.running.get() { + /// If called while another `access` closure is running (re-entrant or from + /// another thread), `f` is queued and run after the current closure and + /// any other queued closures finish. + /// + /// The closure must be `Send` so it can be moved into the queue when called + /// from another thread. + pub fn access(&self, f: impl FnOnce(&mut T) + Send + 'static) { + // If already running (this thread re-entrant or another thread), enqueue. + if self.running.swap(true, Ordering::Acquire) { self.queue.lock().unwrap().push_back(Box::new(f)); - return; } - // first entrant = executor - self.running.set(true); + // We are the executor: hold the write lock for the whole run + drain. + let mut guard = self.value.write().unwrap(); + f(&mut *guard); - // call current - f(self.access_mut()); - - // drain queued re-entrant calls while let Some(job) = self.queue.lock().unwrap().pop_front() { - let value = self.access_mut(); - job(value); + job(&mut *guard); } - self.running.set(false); - } - - /// Returns a mutable reference to the inner value. - /// - /// # Safety - /// Only call this while you are inside an `access` closure (or while draining - /// the queue). Otherwise you may create aliasing mutable references. - pub fn access_mut(&self) -> &mut T { - unsafe { &mut *self.value.get() } + drop(guard); + self.running.store(false, Ordering::Release); } /// Returns an immutable reference to the inner value. /// - /// Safe to call at any time; no closure is running when you use this for - /// read-only access. - pub fn access_ref(&self) -> &T { - unsafe { &*self.value.get() } + /// Blocks if an `access` closure is currently running. Safe to call from + /// any thread. + pub fn access_ref(&self) -> std::sync::RwLockReadGuard<'_, T> { + self.value.read().unwrap() } } +#[cfg(test)] mod tests { + use std::sync::Arc; + use std::thread; + #[test] fn test_access() { - let value = std::sync::Arc::new(crate::AccessCell::new(0)); + let value = Arc::new(crate::AccessCell::new(0)); value.access({ let value = value.clone(); @@ -95,4 +99,29 @@ mod tests { assert_eq!(*value.access_ref(), 10); } + + #[test] + fn test_access_multithreaded() { + let cell = Arc::new(crate::AccessCell::new(0i32)); + let num_threads = 8; + let increments_per_thread = 100; + + let handles: Vec<_> = (0..num_threads) + .map(|_| { + let cell = Arc::clone(&cell); + thread::spawn(move || { + for _ in 0..increments_per_thread { + cell.access(|v| *v += 1); + } + }) + }) + .collect(); + + for h in handles { + h.join().unwrap(); + } + + let expected = num_threads * increments_per_thread; + assert_eq!(*cell.access_ref(), expected); + } }