moirai_async/sync/
rwlock.rs1#![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
22pub struct RwLock<T> {
24 data: UnsafeCell<T>,
25 state: Mutex<RwLockState>,
26}
27
28unsafe impl<T: Send + Sync> Sync for RwLock<T> {}
33unsafe impl<T: Send> Send for RwLock<T> {}
35
36struct RwLockState {
37 readers: usize,
38 writer: bool,
39 read_waiters: WaitQueue<()>,
42 write_waiters: WaitQueue<()>,
43}
44
45impl RwLockState {
46 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 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 pub fn read(&self) -> RwLockReadFuture<'_, T> {
71 RwLockReadFuture {
72 lock: self,
73 id: None,
74 }
75 }
76
77 pub fn write(&self) -> RwLockWriteFuture<'_, T> {
79 RwLockWriteFuture {
80 lock: self,
81 id: None,
82 }
83 }
84
85 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 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 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
143pub 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 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 WaiterPoll::NotRegistered => {}
165 }
166 }
167
168 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 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
198pub 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 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 WaiterPoll::NotRegistered => {}
220 }
221 }
222
223 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 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
253pub 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 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
273pub 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 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 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}