1use crate::sync::{ResourceSignature, ResourceType};
2use shuttle_engine::future::batch_semaphore::{BatchSemaphore, Fairness};
3use shuttle_engine::runtime::execution::ExecutionState;
4use shuttle_engine::runtime::task::{TaskId, TaskSet};
5use shuttle_engine::runtime::thread;
6use std::cell::RefCell;
7use std::fmt::{Debug, Display};
8use std::ops::{Deref, DerefMut};
9use std::panic::{RefUnwindSafe, UnwindSafe};
10use std::sync::{LockResult, PoisonError, TryLockError, TryLockResult};
11use tracing::trace;
12
13const MAX_READS: usize = (u32::MAX >> 3) as usize;
16
17pub struct RwLock<T: ?Sized> {
23 state: RefCell<RwLockState>,
24 semaphore: BatchSemaphore,
25 inner: std::sync::RwLock<T>,
26}
27
28#[derive(Debug)]
29struct RwLockState {
30 holder: RwLockHolder,
31}
32
33#[derive(PartialEq, Eq, Debug)]
34enum RwLockHolder {
35 Read(TaskSet),
36 Write(TaskId),
37 None,
38}
39
40#[derive(PartialEq, Eq, Debug, Clone, Copy)]
41enum RwLockType {
42 Read,
43 Write,
44}
45
46impl RwLockType {
47 fn num_permits(&self) -> usize {
49 match self {
50 Self::Read => 1,
51 Self::Write => MAX_READS,
52 }
53 }
54}
55
56impl<T> RwLock<T> {
57 #[track_caller]
59 pub const fn new(value: T) -> Self {
60 let state = RwLockState {
61 holder: RwLockHolder::None,
62 };
63
64 Self {
65 inner: std::sync::RwLock::new(value),
66 semaphore: BatchSemaphore::const_new_with_signature(
67 MAX_READS,
68 Fairness::Unfair,
69 ResourceSignature::new_const(ResourceType::RwLock),
70 ),
71 state: RefCell::new(state),
72 }
73 }
74}
75
76impl<T: ?Sized> RwLock<T> {
77 pub fn read(&self) -> LockResult<RwLockReadGuard<'_, T>> {
80 self.lock(RwLockType::Read);
81
82 match self.inner.try_read() {
83 Ok(guard) => Ok(RwLockReadGuard {
84 inner: Some(guard),
85 rwlock: self,
86 me: ExecutionState::me(),
87 }),
88 Err(TryLockError::Poisoned(err)) => Err(PoisonError::new(RwLockReadGuard {
89 inner: Some(err.into_inner()),
90 rwlock: self,
91 me: ExecutionState::me(),
92 })),
93 Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
94 }
95 }
96
97 pub fn write(&self) -> LockResult<RwLockWriteGuard<'_, T>> {
100 self.lock(RwLockType::Write);
101
102 match self.inner.try_write() {
103 Ok(guard) => Ok(RwLockWriteGuard {
104 inner: Some(guard),
105 rwlock: self,
106 me: ExecutionState::me(),
107 }),
108 Err(TryLockError::Poisoned(err)) => Err(PoisonError::new(RwLockWriteGuard {
109 inner: Some(err.into_inner()),
110 rwlock: self,
111 me: ExecutionState::me(),
112 })),
113 Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
114 }
115 }
116
117 pub fn try_read(&self) -> TryLockResult<RwLockReadGuard<'_, T>> {
125 if self.try_lock(RwLockType::Read) {
126 match self.inner.try_read() {
127 Ok(guard) => Ok(RwLockReadGuard {
128 inner: Some(guard),
129 rwlock: self,
130 me: ExecutionState::me(),
131 }),
132 Err(TryLockError::Poisoned(err)) => Err(TryLockError::Poisoned(PoisonError::new(RwLockReadGuard {
133 inner: Some(err.into_inner()),
134 rwlock: self,
135 me: ExecutionState::me(),
136 }))),
137 Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
138 }
139 } else {
140 Err(TryLockError::WouldBlock)
141 }
142 }
143
144 pub fn try_write(&self) -> TryLockResult<RwLockWriteGuard<'_, T>> {
149 if self.try_lock(RwLockType::Write) {
150 match self.inner.try_write() {
151 Ok(guard) => Ok(RwLockWriteGuard {
152 inner: Some(guard),
153 rwlock: self,
154 me: ExecutionState::me(),
155 }),
156 Err(TryLockError::Poisoned(err)) => Err(TryLockError::Poisoned(PoisonError::new(RwLockWriteGuard {
157 inner: Some(err.into_inner()),
158 rwlock: self,
159 me: ExecutionState::me(),
160 }))),
161 Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
162 }
163 } else {
164 Err(TryLockError::WouldBlock)
165 }
166 }
167
168 #[inline]
173 pub fn get_mut(&mut self) -> LockResult<&mut T> {
174 self.inner.get_mut()
175 }
176
177 pub fn into_inner(self) -> LockResult<T>
179 where
180 T: Sized,
181 {
182 let state = self.state.borrow();
183 assert_eq!(state.holder, RwLockHolder::None);
184
185 self.semaphore.try_acquire(MAX_READS).unwrap();
187
188 self.inner.into_inner()
189 }
190
191 fn lock(&self, typ: RwLockType) {
193 let me = ExecutionState::me();
194
195 let mut state = self.state.borrow_mut();
196 trace!(
197 holder = ?state.holder,
198 semaphore = ?self.semaphore,
199 "acquiring {:?} lock on rwlock {:p}",
200 typ,
201 self,
202 );
203 drop(state);
204
205 if !self.semaphore.is_closed() {
206 state = self.state.borrow_mut();
208 assert!(
209 match &state.holder {
210 RwLockHolder::Write(writer) => *writer != me,
211 RwLockHolder::Read(readers) => !readers.contains(me),
212 RwLockHolder::None => true,
213 },
214 "deadlock! task {me:?} tried to acquire a RwLock it already holds"
215 );
216 drop(state);
217
218 self.semaphore.acquire_blocking(typ.num_permits()).unwrap();
219 } else {
220 thread::switch();
222 }
223
224 state = self.state.borrow_mut();
225 match (typ, &mut state.holder) {
226 (RwLockType::Write, RwLockHolder::None) => {
227 state.holder = RwLockHolder::Write(me);
228 }
229 (RwLockType::Read, RwLockHolder::None) => {
230 let mut readers = TaskSet::new();
231 readers.insert(me);
232 state.holder = RwLockHolder::Read(readers);
233 }
234 (RwLockType::Read, RwLockHolder::Read(readers)) => {
235 assert!(readers.insert(me));
236 }
237 _ => {
238 panic!(
239 "resumed a waiting {:?} thread while the lock was in state {:?}",
240 typ, state.holder
241 );
242 }
243 }
244 trace!(
245 holder = ?state.holder,
246 semaphore = ?self.semaphore,
247 "acquired {:?} lock on rwlock {:p}",
248 typ,
249 self
250 );
251 drop(state);
252 }
253
254 fn try_lock(&self, typ: RwLockType) -> bool {
257 let me = ExecutionState::me();
258
259 let mut state = self.state.borrow_mut();
260 trace!(
261 holder = ?state.holder,
262 semaphore = ?self.semaphore,
263 "trying to acquire {:?} lock on rwlock {:p}",
264 typ,
265 self,
266 );
267 drop(state);
268
269 let mut acquired = self.semaphore.try_acquire(typ.num_permits()).is_ok();
271 if acquired {
272 state = self.state.borrow_mut();
273 match (typ, &mut state.holder) {
274 (RwLockType::Write, RwLockHolder::None) => {
275 state.holder = RwLockHolder::Write(me);
276 }
277 (RwLockType::Read, RwLockHolder::None) => {
278 let mut readers = TaskSet::new();
279 readers.insert(me);
280 state.holder = RwLockHolder::Read(readers);
281 }
282 (RwLockType::Read, RwLockHolder::Read(readers)) => {
283 acquired = readers.insert(me);
286 }
287 _ => (),
288 };
289 drop(state);
290 }
291
292 trace!(
293 "{} {:?} lock on rwlock {:p}",
294 if acquired { "acquired" } else { "failed to acquire" },
295 typ,
296 self,
297 );
298
299 acquired
300 }
301
302 #[inline]
304 pub fn clear_poison(&self) {
305 self.inner.clear_poison();
306 }
307}
308
309unsafe impl<T: Send + ?Sized> Send for RwLock<T> {}
314unsafe impl<T: Send + ?Sized> Sync for RwLock<T> {}
315
316impl<T: ?Sized> UnwindSafe for RwLock<T> {}
318impl<T: ?Sized> RefUnwindSafe for RwLock<T> {}
319
320impl<T: Default> Default for RwLock<T> {
321 fn default() -> Self {
322 Self::new(Default::default())
323 }
324}
325
326impl<T: ?Sized + Debug> Debug for RwLock<T> {
327 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
328 Debug::fmt(&self.inner, f)
329 }
330}
331
332pub struct RwLockReadGuard<'a, T: ?Sized> {
334 inner: Option<std::sync::RwLockReadGuard<'a, T>>,
335 rwlock: &'a RwLock<T>,
336 me: TaskId,
337}
338
339impl<T: ?Sized> Deref for RwLockReadGuard<'_, T> {
340 type Target = T;
341
342 fn deref(&self) -> &Self::Target {
343 self.inner.as_ref().unwrap().deref()
344 }
345}
346
347impl<T: Debug> Debug for RwLockReadGuard<'_, T> {
348 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
349 Debug::fmt(&self.inner.as_ref().unwrap(), f)
350 }
351}
352
353impl<T: Display + ?Sized> Display for RwLockReadGuard<'_, T> {
354 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
355 (**self).fmt(f)
356 }
357}
358
359impl<T: ?Sized> Drop for RwLockReadGuard<'_, T> {
360 fn drop(&mut self) {
361 struct Unlock<'g, 'a, T: ?Sized>(&'g mut RwLockReadGuard<'a, T>);
366 impl<T: ?Sized> Drop for Unlock<'_, '_, T> {
367 #[inline]
368 fn drop(&mut self) {
369 let guard = &mut *self.0;
370 guard.inner = None;
371
372 let mut state = guard.rwlock.state.borrow_mut();
373 trace!(
374 holder = ?state.holder,
375 semaphore = ?guard.rwlock.semaphore,
376 "releasing Read lock on rwlock {:p}",
377 guard.rwlock
378 );
379 let RwLockHolder::Read(readers) = &mut state.holder else {
380 panic!("exiting a reader but rwlock is in the wrong state {:?}", state.holder);
381 };
382 assert!(readers.remove(guard.me));
383 if readers.is_empty() {
384 state.holder = RwLockHolder::None;
385 }
386 drop(state);
387 }
388 }
389 let unlock = Unlock(self);
390
391 unlock.0.rwlock.semaphore.release(RwLockType::Read.num_permits());
392 }
393}
394
395pub struct RwLockWriteGuard<'a, T: ?Sized> {
397 inner: Option<std::sync::RwLockWriteGuard<'a, T>>,
398 rwlock: &'a RwLock<T>,
399 me: TaskId,
400}
401
402impl<T: ?Sized> Deref for RwLockWriteGuard<'_, T> {
403 type Target = T;
404
405 fn deref(&self) -> &Self::Target {
406 self.inner.as_ref().unwrap().deref()
407 }
408}
409
410impl<T: ?Sized> DerefMut for RwLockWriteGuard<'_, T> {
411 fn deref_mut(&mut self) -> &mut Self::Target {
412 self.inner.as_mut().unwrap().deref_mut()
413 }
414}
415
416impl<T: Debug> Debug for RwLockWriteGuard<'_, T> {
417 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
418 Debug::fmt(&self.inner.as_ref().unwrap(), f)
419 }
420}
421
422impl<T: Display + ?Sized> Display for RwLockWriteGuard<'_, T> {
423 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
424 (**self).fmt(f)
425 }
426}
427
428impl<T: ?Sized> Drop for RwLockWriteGuard<'_, T> {
429 fn drop(&mut self) {
430 struct Unlock<'g, 'a, T: ?Sized>(&'g mut RwLockWriteGuard<'a, T>);
432 impl<T: ?Sized> Drop for Unlock<'_, '_, T> {
433 #[inline]
434 fn drop(&mut self) {
435 let guard = &mut *self.0;
436 let clear_poison = std::thread::panicking()
438 && !guard.rwlock.inner.is_poisoned()
439 && ExecutionState::unwinding_for_teardown();
440 guard.inner = None;
441 if clear_poison {
442 guard.rwlock.inner.clear_poison();
443 }
444
445 let mut state = guard.rwlock.state.borrow_mut();
446 trace!(
447 holder = ?state.holder,
448 semaphore = ?guard.rwlock.semaphore,
449 "releasing Write lock on rwlock {:p}",
450 guard.rwlock
451 );
452 assert_eq!(state.holder, RwLockHolder::Write(guard.me));
453 state.holder = RwLockHolder::None;
454 drop(state);
455 }
456 }
457 let unlock = Unlock(self);
458
459 unlock.0.rwlock.semaphore.release(RwLockType::Write.num_permits());
460 }
461}
462
463#[cfg(test)]
464mod tests {
465 use super::*;
466
467 #[test]
468 fn unique_resource_signature_rwlock() {
469 shuttle_schedulers::check_random(
470 || {
471 let rwlock1 = RwLock::new(0);
472 let rwlock2 = RwLock::new(0);
473 assert_ne!(rwlock1.semaphore.signature(), rwlock2.semaphore.signature());
474 },
475 1,
476 );
477 }
478}