Thread safe
This commit is contained in:
+66
-37
@@ -3,11 +3,15 @@
|
|||||||
//! When you call [`access`](AccessCell::access) with a closure that itself calls `access`
|
//! 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
|
//! on the same cell, those nested calls are queued and run after the current closure
|
||||||
//! finishes, avoiding deadlock.
|
//! finishes, avoiding deadlock.
|
||||||
|
//!
|
||||||
|
//! This type is thread-safe: it implements [`Send`] and [`Sync`] when `T: Send`.
|
||||||
|
|
||||||
use std::{
|
use std::{
|
||||||
cell::{Cell, UnsafeCell},
|
|
||||||
collections::VecDeque,
|
collections::VecDeque,
|
||||||
sync::Mutex,
|
sync::{
|
||||||
|
atomic::{AtomicBool, Ordering},
|
||||||
|
Mutex, RwLock,
|
||||||
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
/// A cell holding a value of type `T` with re-entrant mutable access.
|
/// 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
|
/// The first caller of [`access`](AccessCell::access) runs immediately; any further
|
||||||
/// calls to `access` from within that closure (or its queue) are enqueued and
|
/// calls to `access` from within that closure (or its queue) are enqueued and
|
||||||
/// executed in order after the current work completes.
|
/// executed in order after the current work completes.
|
||||||
|
///
|
||||||
|
/// Thread-safe: safe to share across threads and use from multiple threads.
|
||||||
pub struct AccessCell<T> {
|
pub struct AccessCell<T> {
|
||||||
value: UnsafeCell<T>,
|
value: RwLock<T>,
|
||||||
running: Cell<bool>,
|
running: AtomicBool,
|
||||||
queue: Mutex<VecDeque<Box<dyn FnOnce(&mut T)>>>,
|
queue: Mutex<VecDeque<Box<dyn FnOnce(&mut T) + Send>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Safe: only one thread holds the write guard at a time (enforced by `running` + queue).
|
||||||
|
unsafe impl<T: Send> Send for AccessCell<T> {}
|
||||||
|
unsafe impl<T: Send> Sync for AccessCell<T> {}
|
||||||
|
|
||||||
impl<T> AccessCell<T> {
|
impl<T> AccessCell<T> {
|
||||||
/// Creates a new `AccessCell` wrapping `value`.
|
/// Creates a new `AccessCell` wrapping `value`.
|
||||||
pub fn new(value: T) -> Self {
|
pub fn new(value: T) -> Self {
|
||||||
Self {
|
Self {
|
||||||
value: UnsafeCell::new(value),
|
value: RwLock::new(value),
|
||||||
running: Cell::new(false),
|
running: AtomicBool::new(false),
|
||||||
queue: Mutex::new(VecDeque::new()),
|
queue: Mutex::new(VecDeque::new()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Runs `f` with exclusive mutable access to the inner value.
|
/// Runs `f` with exclusive mutable access to the inner value.
|
||||||
///
|
///
|
||||||
/// If called while another `access` closure is running (re-entrant call),
|
/// If called while another `access` closure is running (re-entrant or from
|
||||||
/// `f` is queued and run after the current closure and any other queued
|
/// another thread), `f` is queued and run after the current closure and
|
||||||
/// closures finish.
|
/// any other queued closures finish.
|
||||||
pub fn access(&self, f: impl FnOnce(&mut T) + 'static) {
|
///
|
||||||
// already inside → enqueue
|
/// The closure must be `Send` so it can be moved into the queue when called
|
||||||
if self.running.get() {
|
/// 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));
|
self.queue.lock().unwrap().push_back(Box::new(f));
|
||||||
|
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
// first entrant = executor
|
// We are the executor: hold the write lock for the whole run + drain.
|
||||||
self.running.set(true);
|
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() {
|
while let Some(job) = self.queue.lock().unwrap().pop_front() {
|
||||||
let value = self.access_mut();
|
job(&mut *guard);
|
||||||
job(value);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
self.running.set(false);
|
drop(guard);
|
||||||
}
|
self.running.store(false, Ordering::Release);
|
||||||
|
|
||||||
/// 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() }
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Returns an immutable reference to the inner value.
|
/// Returns an immutable reference to the inner value.
|
||||||
///
|
///
|
||||||
/// Safe to call at any time; no closure is running when you use this for
|
/// Blocks if an `access` closure is currently running. Safe to call from
|
||||||
/// read-only access.
|
/// any thread.
|
||||||
pub fn access_ref(&self) -> &T {
|
pub fn access_ref(&self) -> std::sync::RwLockReadGuard<'_, T> {
|
||||||
unsafe { &*self.value.get() }
|
self.value.read().unwrap()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
|
use std::sync::Arc;
|
||||||
|
use std::thread;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_access() {
|
fn test_access() {
|
||||||
let value = std::sync::Arc::new(crate::AccessCell::new(0));
|
let value = Arc::new(crate::AccessCell::new(0));
|
||||||
|
|
||||||
value.access({
|
value.access({
|
||||||
let value = value.clone();
|
let value = value.clone();
|
||||||
@@ -95,4 +99,29 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(*value.access_ref(), 10);
|
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);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user