1use core::cell::UnsafeCell;
22use core::marker::PhantomData;
23use core::ops::{Deref, DerefMut};
24use core::sync::atomic::{AtomicUsize, Ordering};
25
26use crate::{Mutex, Semaphore};
27
28struct RawRwLock {
31 state: AtomicUsize,
32 mutex: Mutex,
33 semaphore: Semaphore,
34}
35
36const COUNT_BITS: u32 = (usize::BITS - 1) / 2;
45const COUNT_MAX: usize = (1usize << COUNT_BITS) - 1;
46
47const IS_WRITING: usize = 1;
48const WRITER: usize = 1 << 1;
49const READER: usize = 1 << (1 + COUNT_BITS);
50const WRITER_MASK: usize = COUNT_MAX << WRITER.trailing_zeros();
51const READER_MASK: usize = COUNT_MAX << READER.trailing_zeros();
52
53impl RawRwLock {
54 const fn new() -> Self {
55 Self {
56 state: AtomicUsize::new(0),
57 mutex: Mutex::new(),
58 semaphore: Semaphore::new(),
59 }
60 }
61
62 fn try_lock(&self) -> bool {
63 if self.mutex.try_lock() {
64 let state = self.state.load(Ordering::SeqCst);
65 if state & READER_MASK == 0 {
66 let _ = self.state.fetch_or(IS_WRITING, Ordering::SeqCst);
67 return true;
68 }
69
70 self.mutex.unlock();
71 }
72
73 false
74 }
75
76 fn lock(&self) {
77 let _ = self.state.fetch_add(WRITER, Ordering::SeqCst);
78 self.mutex.lock();
79
80 let state = self
83 .state
84 .fetch_add(IS_WRITING.wrapping_sub(WRITER), Ordering::SeqCst);
85 if state & READER_MASK != 0 {
86 self.semaphore.wait();
87 }
88 }
89
90 fn unlock(&self) {
91 let _ = self.state.fetch_and(!IS_WRITING, Ordering::SeqCst);
92 self.mutex.unlock();
93 }
94
95 fn try_lock_shared(&self) -> bool {
96 let state = self.state.load(Ordering::SeqCst);
97 if state & (IS_WRITING | WRITER_MASK) == 0 {
98 if self
100 .state
101 .compare_exchange(state, state + READER, Ordering::SeqCst, Ordering::SeqCst)
102 .is_ok()
103 {
104 return true;
105 }
106 }
107
108 if self.mutex.try_lock() {
109 let _ = self.state.fetch_add(READER, Ordering::SeqCst);
110 self.mutex.unlock();
111 return true;
112 }
113
114 false
115 }
116
117 fn lock_shared(&self) {
118 let mut state = self.state.load(Ordering::SeqCst);
119 while state & (IS_WRITING | WRITER_MASK) == 0 {
120 match self.state.compare_exchange_weak(
122 state,
123 state + READER,
124 Ordering::SeqCst,
125 Ordering::SeqCst,
126 ) {
127 Ok(_) => return,
128 Err(s) => state = s,
129 }
130 }
131
132 self.mutex.lock();
133 let _ = self.state.fetch_add(READER, Ordering::SeqCst);
134 self.mutex.unlock();
135 }
136
137 fn unlock_shared(&self) {
138 let state = self.state.fetch_sub(READER, Ordering::SeqCst);
139
140 if (state & READER_MASK == READER) && (state & IS_WRITING != 0) {
141 self.semaphore.post();
142 }
143 }
144}
145
146pub struct RwLock<T> {
150 raw: RawRwLock,
151 value: UnsafeCell<T>,
152}
153
154unsafe impl<T: Send> Send for RwLock<T> {}
158unsafe impl<T: Send + Sync> Sync for RwLock<T> {}
162
163impl<T: Default> Default for RwLock<T> {
164 fn default() -> Self {
165 Self::new(T::default())
166 }
167}
168
169impl<T> RwLock<T> {
170 pub const fn new(value: T) -> Self {
173 Self {
174 raw: RawRwLock::new(),
175 value: UnsafeCell::new(value),
176 }
177 }
178
179 #[inline]
182 pub fn read(&self) -> RwLockReadGuard<'_, T> {
183 self.raw.lock_shared();
184 RwLockReadGuard {
185 lock: self,
186 _not_send: PhantomData,
187 }
188 }
189
190 #[inline]
193 pub fn write(&self) -> RwLockWriteGuard<'_, T> {
194 self.raw.lock();
195 RwLockWriteGuard {
196 lock: self,
197 _not_send: PhantomData,
198 }
199 }
200
201 #[inline]
203 pub fn try_read(&self) -> Option<RwLockReadGuard<'_, T>> {
204 if self.raw.try_lock_shared() {
205 Some(RwLockReadGuard {
206 lock: self,
207 _not_send: PhantomData,
208 })
209 } else {
210 None
211 }
212 }
213
214 #[inline]
216 pub fn try_write(&self) -> Option<RwLockWriteGuard<'_, T>> {
217 if self.raw.try_lock() {
218 Some(RwLockWriteGuard {
219 lock: self,
220 _not_send: PhantomData,
221 })
222 } else {
223 None
224 }
225 }
226
227 #[inline]
230 pub fn get_mut(&mut self) -> &mut T {
231 self.value.get_mut()
232 }
233
234 #[inline]
236 pub fn into_inner(self) -> T {
237 self.value.into_inner()
238 }
239}
240
241pub struct RwLockReadGuard<'a, T> {
247 lock: &'a RwLock<T>,
248 _not_send: PhantomData<*const ()>,
249}
250
251impl<'a, T> Deref for RwLockReadGuard<'a, T> {
252 type Target = T;
253 #[inline]
254 fn deref(&self) -> &T {
255 unsafe { &*self.lock.value.get() }
257 }
258}
259
260impl<'a, T> Drop for RwLockReadGuard<'a, T> {
261 #[inline]
262 fn drop(&mut self) {
263 self.lock.raw.unlock_shared();
264 }
265}
266
267pub struct RwLockWriteGuard<'a, T> {
272 lock: &'a RwLock<T>,
273 _not_send: PhantomData<*const ()>,
274}
275
276impl<'a, T> Deref for RwLockWriteGuard<'a, T> {
277 type Target = T;
278 #[inline]
279 fn deref(&self) -> &T {
280 unsafe { &*self.lock.value.get() }
282 }
283}
284
285impl<'a, T> DerefMut for RwLockWriteGuard<'a, T> {
286 #[inline]
287 fn deref_mut(&mut self) -> &mut T {
288 unsafe { &mut *self.lock.value.get() }
290 }
291}
292
293impl<'a, T> Drop for RwLockWriteGuard<'a, T> {
294 #[inline]
295 fn drop(&mut self) {
296 self.lock.raw.unlock();
297 }
298}
299
300#[cfg(test)]
301mod tests {
302 use super::*;
303
304 #[test]
305 fn smoke() {
306 let rwl = RwLock::new(0u32);
307
308 {
309 let mut w = rwl.write();
310 assert!(rwl.try_write().is_none());
311 assert!(rwl.try_read().is_none());
312 *w = 1;
313 }
314
315 {
316 let w = rwl.try_write().unwrap();
317 assert!(rwl.try_write().is_none());
318 assert!(rwl.try_read().is_none());
319 drop(w);
320 }
321
322 {
323 let r1 = rwl.read();
324 assert!(rwl.try_write().is_none());
325 let r2 = rwl.try_read().unwrap();
326 assert_eq!(*r1, 1);
327 assert_eq!(*r2, 1);
328 }
329
330 {
331 let r1 = rwl.try_read().unwrap();
332 assert!(rwl.try_write().is_none());
333 let r2 = rwl.try_read().unwrap();
334 drop((r1, r2));
335 }
336
337 let _w = rwl.write();
338 }
339
340 #[test]
341 fn raw_internal_state() {
342 let raw = RawRwLock::new();
345 raw.lock();
346 raw.unlock();
347 assert_eq!(raw.state.load(Ordering::SeqCst), 0);
348 }
349}
350
351