Skip to main content

moirai_async/sync/
rwlock.rs

1//! Async-aware RwLock for concurrent read/exclusive write access
2//!
3//! Provides an async-compatible RwLock that allows multiple concurrent readers
4//! or a single writer, following SLAP principle design. Waiter-queue mechanics
5//! live in `WaitQueue`; this module keeps only the reader/writer admission
6//! predicates (writer preference for pending writers, reader-batch grants on
7//! writer release) and the lock-restoration policy for cancelled futures.
8
9#![expect(
10    clippy::unwrap_used,
11    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
12)]
13
14use std::cell::UnsafeCell;
15use std::future::Future;
16use std::pin::Pin;
17use std::sync::Mutex;
18use std::task::{Context, Poll, Waker};
19
20use crate::sync::wait_queue::{WaitQueue, WaiterPoll};
21
22/// Async-aware RwLock
23pub struct RwLock<T> {
24    data: UnsafeCell<T>,
25    state: Mutex<RwLockState>,
26}
27
28// SAFETY: access to `data` is mediated exclusively by the guard types, whose
29// issuance is serialized through `state` (readers-shared XOR writer-exclusive).
30// `T: Send + Sync` is required so shared references handed out to concurrent
31// readers on other threads are sound.
32unsafe impl<T: Send + Sync> Sync for RwLock<T> {}
33// SAFETY: moving the lock moves `data` by value; only `T: Send` is required.
34unsafe impl<T: Send> Send for RwLock<T> {}
35
36struct RwLockState {
37    readers: usize,
38    writer: bool,
39    /// Reader and writer waiters in separate FIFO queues; a grant hands the
40    /// lock directly to the waiter (`()` payload — the grant is the lock).
41    read_waiters: WaitQueue<()>,
42    write_waiters: WaitQueue<()>,
43}
44
45impl RwLockState {
46    /// Grant the lock to the oldest ungranted writer, marking `writer`, and
47    /// return its waker to wake. Returns `None` if no writer is waiting.
48    fn grant_oldest_writer(&mut self) -> Option<Waker> {
49        let waker = self.write_waiters.grant_oldest(())?;
50        self.writer = true;
51        Some(waker)
52    }
53}
54
55impl<T> RwLock<T> {
56    /// Create a new async RwLock
57    pub fn new(data: T) -> Self {
58        Self {
59            data: UnsafeCell::new(data),
60            state: Mutex::new(RwLockState {
61                readers: 0,
62                writer: false,
63                read_waiters: WaitQueue::new(),
64                write_waiters: WaitQueue::new(),
65            }),
66        }
67    }
68
69    /// Acquire a read lock asynchronously
70    pub fn read(&self) -> RwLockReadFuture<'_, T> {
71        RwLockReadFuture {
72            lock: self,
73            id: None,
74        }
75    }
76
77    /// Acquire a write lock asynchronously
78    pub fn write(&self) -> RwLockWriteFuture<'_, T> {
79        RwLockWriteFuture {
80            lock: self,
81            id: None,
82        }
83    }
84
85    /// Try to acquire a read lock immediately
86    pub fn try_read(&self) -> Option<RwLockReadGuard<'_, T>> {
87        let mut state = self.state.lock().unwrap();
88        if !state.writer && state.write_waiters.is_empty() {
89            state.readers += 1;
90            Some(RwLockReadGuard { lock: self })
91        } else {
92            None
93        }
94    }
95
96    /// Try to acquire a write lock immediately
97    pub fn try_write(&self) -> Option<RwLockWriteGuard<'_, T>> {
98        let mut state = self.state.lock().unwrap();
99        if state.readers == 0 && !state.writer {
100            state.writer = true;
101            Some(RwLockWriteGuard { lock: self })
102        } else {
103            None
104        }
105    }
106
107    fn release_read(&self) {
108        let mut state = self.state.lock().unwrap();
109        state.readers -= 1;
110        if state.readers == 0 {
111            let waker = state.grant_oldest_writer();
112            drop(state);
113            if let Some(w) = waker {
114                w.wake();
115            }
116        }
117    }
118
119    fn release_write(&self) {
120        let mut state = self.state.lock().unwrap();
121        state.writer = false;
122
123        // Prefer waking every pending reader (reader batch); only if there are
124        // none, hand the lock to the oldest waiting writer.
125        let reader_wakers = state.read_waiters.grant_all(());
126
127        if !reader_wakers.is_empty() {
128            state.readers += reader_wakers.len();
129            drop(state);
130            for waker in reader_wakers {
131                waker.wake();
132            }
133        } else {
134            let waker = state.grant_oldest_writer();
135            drop(state);
136            if let Some(w) = waker {
137                w.wake();
138            }
139        }
140    }
141}
142
143/// Future for async read lock acquisition
144pub struct RwLockReadFuture<'a, T> {
145    lock: &'a RwLock<T>,
146    id: Option<u64>,
147}
148
149impl<'a, T> Future for RwLockReadFuture<'a, T> {
150    type Output = RwLockReadGuard<'a, T>;
151
152    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
153        let mut state = self.lock.state.lock().unwrap();
154
155        // 1. Check if we were already registered and have been granted the lock
156        if let Some(id) = self.id {
157            match state.read_waiters.poll_waiter(id, cx.waker()) {
158                WaiterPoll::Granted(()) => {
159                    self.id = None;
160                    return Poll::Ready(RwLockReadGuard { lock: self.lock });
161                }
162                WaiterPoll::Pending => return Poll::Pending,
163                // registration lost; fall through
164                WaiterPoll::NotRegistered => {}
165            }
166        }
167
168        // 2. Try to acquire the read lock
169        if !state.writer && state.write_waiters.is_empty() {
170            state.readers += 1;
171            if let Some(id) = self.id.take() {
172                let _removed_grant = state.read_waiters.deregister(id);
173            }
174            return Poll::Ready(RwLockReadGuard { lock: self.lock });
175        }
176
177        // 3. Register as a reader waiter
178        if self.id.is_none() {
179            self.id = Some(state.read_waiters.register(cx.waker().clone()));
180        }
181
182        Poll::Pending
183    }
184}
185
186impl<'a, T> Drop for RwLockReadFuture<'a, T> {
187    fn drop(&mut self) {
188        if let Some(id) = self.id
189            && let Ok(mut state) = self.lock.state.lock()
190            && state.read_waiters.deregister(id).is_some()
191        {
192            drop(state);
193            self.lock.release_read();
194        }
195    }
196}
197
198/// Future for async write lock acquisition
199pub struct RwLockWriteFuture<'a, T> {
200    lock: &'a RwLock<T>,
201    id: Option<u64>,
202}
203
204impl<'a, T> Future for RwLockWriteFuture<'a, T> {
205    type Output = RwLockWriteGuard<'a, T>;
206
207    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
208        let mut state = self.lock.state.lock().unwrap();
209
210        // 1. Check if we were already registered and have been granted the lock
211        if let Some(id) = self.id {
212            match state.write_waiters.poll_waiter(id, cx.waker()) {
213                WaiterPoll::Granted(()) => {
214                    self.id = None;
215                    return Poll::Ready(RwLockWriteGuard { lock: self.lock });
216                }
217                WaiterPoll::Pending => return Poll::Pending,
218                // registration lost; fall through
219                WaiterPoll::NotRegistered => {}
220            }
221        }
222
223        // 2. Try to acquire the write lock
224        if state.readers == 0 && !state.writer {
225            state.writer = true;
226            if let Some(id) = self.id.take() {
227                let _removed_grant = state.write_waiters.deregister(id);
228            }
229            return Poll::Ready(RwLockWriteGuard { lock: self.lock });
230        }
231
232        // 3. Register as a writer waiter
233        if self.id.is_none() {
234            self.id = Some(state.write_waiters.register(cx.waker().clone()));
235        }
236
237        Poll::Pending
238    }
239}
240
241impl<'a, T> Drop for RwLockWriteFuture<'a, T> {
242    fn drop(&mut self) {
243        if let Some(id) = self.id
244            && let Ok(mut state) = self.lock.state.lock()
245            && state.write_waiters.deregister(id).is_some()
246        {
247            drop(state);
248            self.lock.release_write();
249        }
250    }
251}
252
253/// Guard for RwLock read access
254pub struct RwLockReadGuard<'a, T> {
255    lock: &'a RwLock<T>,
256}
257
258impl<'a, T> std::ops::Deref for RwLockReadGuard<'a, T> {
259    type Target = T;
260    fn deref(&self) -> &Self::Target {
261        // SAFETY: guard existence implies a held read lock (`readers > 0`,
262        // `writer == false`), so shared access to `data` is sound.
263        unsafe { &*self.lock.data.get() }
264    }
265}
266
267impl<'a, T> Drop for RwLockReadGuard<'a, T> {
268    fn drop(&mut self) {
269        self.lock.release_read();
270    }
271}
272
273/// Guard for RwLock write access
274pub struct RwLockWriteGuard<'a, T> {
275    lock: &'a RwLock<T>,
276}
277
278impl<'a, T> std::ops::Deref for RwLockWriteGuard<'a, T> {
279    type Target = T;
280    fn deref(&self) -> &Self::Target {
281        // SAFETY: guard existence implies the held exclusive write lock.
282        unsafe { &*self.lock.data.get() }
283    }
284}
285
286impl<'a, T> std::ops::DerefMut for RwLockWriteGuard<'a, T> {
287    fn deref_mut(&mut self) -> &mut Self::Target {
288        // SAFETY: guard existence implies the held exclusive write lock, and
289        // `&mut self` guarantees this is the sole live reference through it.
290        unsafe { &mut *self.lock.data.get() }
291    }
292}
293
294impl<'a, T> Drop for RwLockWriteGuard<'a, T> {
295    fn drop(&mut self) {
296        self.lock.release_write();
297    }
298}
299
300#[cfg(test)]
301mod tests {
302    use super::RwLock;
303    use std::future::Future;
304    use std::pin::Pin;
305    use std::task::{Context, Poll, Waker};
306
307    fn poll_future<F>(future: &mut F) -> Poll<F::Output>
308    where
309        F: Future + Unpin,
310    {
311        let mut context = Context::from_waker(Waker::noop());
312        Pin::new(future).poll(&mut context)
313    }
314
315    #[test]
316    fn last_reader_release_grants_first_waiting_writer() {
317        let lock = RwLock::new(5_u32);
318        let reader = lock.try_read().expect("read lock must be acquired");
319        let mut writer = lock.write();
320
321        assert!(matches!(poll_future(&mut writer), Poll::Pending));
322
323        drop(reader);
324
325        match poll_future(&mut writer) {
326            Poll::Ready(mut guard) => {
327                *guard += 7;
328            }
329            Poll::Pending => panic!("writer waiter must be granted after final reader release"),
330        }
331
332        let reader = lock
333            .try_read()
334            .expect("read lock must be acquired after writer release");
335        assert_eq!(*reader, 12);
336    }
337
338    #[test]
339    fn writer_release_grants_all_registered_readers() {
340        let lock = RwLock::new(11_u32);
341        let writer = lock.try_write().expect("write lock must be acquired");
342        let mut first_reader = lock.read();
343        let mut second_reader = lock.read();
344
345        assert!(matches!(poll_future(&mut first_reader), Poll::Pending));
346        assert!(matches!(poll_future(&mut second_reader), Poll::Pending));
347
348        drop(writer);
349
350        let first_guard = match poll_future(&mut first_reader) {
351            Poll::Ready(guard) => guard,
352            Poll::Pending => panic!("first reader waiter must be granted after writer release"),
353        };
354        let second_guard = match poll_future(&mut second_reader) {
355            Poll::Ready(guard) => guard,
356            Poll::Pending => panic!("second reader waiter must be granted after writer release"),
357        };
358
359        assert_eq!(*first_guard, 11);
360        assert_eq!(*second_guard, 11);
361        assert!(
362            lock.try_write().is_none(),
363            "active granted readers must exclude writers"
364        );
365
366        drop(first_guard);
367        drop(second_guard);
368
369        let mut writer = lock
370            .try_write()
371            .expect("write lock must be acquired after readers release");
372        *writer = 19;
373        drop(writer);
374
375        let reader = lock
376            .try_read()
377            .expect("read lock must be acquired after writer release");
378        assert_eq!(*reader, 19);
379    }
380}