1use crate::task::CurrentTask;
6use starnix_uapi::errors::Errno;
7
8use core::marker::PhantomData;
9
10use starnix_sync::{InterruptibleEvent, LockDepMutex, LockLevel, RwQueueInnerLock};
11use std::collections::VecDeque;
12use std::sync::Arc;
13
14use lock_api as _;
15
16#[cfg(any(test, debug_assertions))]
17use lock_api::RawRwLock;
18
19#[derive(Debug)]
20pub struct RwQueue<L> {
21 inner: LockDepMutex<RwQueueInner, RwQueueInnerLock>,
22 _phantom: PhantomData<L>,
23
24 #[cfg(any(test, debug_assertions))]
26 tracer: tracer::MutexTracer,
27}
28
29impl<L> RwQueue<L> {
30 fn unlock_read(&self) {
31 self.inner.lock().unlock_read();
32
33 #[allow(
34 clippy::undocumented_unsafe_blocks,
35 reason = "Force documented unsafe blocks in Starnix"
36 )]
37 #[cfg(any(test, debug_assertions))]
38 unsafe {
39 self.tracer.unlock_shared();
40 }
41 }
42
43 fn unlock_write(&self) {
44 self.inner.lock().unlock_write();
45
46 #[allow(
47 clippy::undocumented_unsafe_blocks,
48 reason = "Force documented unsafe blocks in Starnix"
49 )]
50 #[cfg(any(test, debug_assertions))]
51 unsafe {
52 self.tracer.unlock_exclusive();
53 }
54 }
55}
56
57impl<L: LockLevel> RwQueue<L> {
58 pub fn read<'a>(
59 &'a self,
60 current_task: &CurrentTask,
61 ) -> Result<RwQueueReadGuard<'a, L>, Errno> {
62 #[cfg(any(test, debug_assertions))]
65 self.tracer.lock_shared();
66
67 let mut inner = self.inner.lock();
68
69 if !inner.try_read() {
70 let event = InterruptibleEvent::new();
71 let guard = event.begin_wait();
72
73 inner.waiters.push_back(Waiter::Reader(event.clone()));
74
75 std::mem::drop(inner);
76
77 current_task.block_until(guard, zx::MonotonicInstant::INFINITE).map_err(|e| {
78 self.inner.lock().remove_waiter(&event);
79 e
80 })?;
81 }
82 Ok(RwQueueReadGuard { queue: self })
83 }
84
85 pub fn write<'a>(
86 &'a self,
87 current_task: &CurrentTask,
88 ) -> Result<RwQueueWriteGuard<'a, L>, Errno> {
89 #[cfg(any(test, debug_assertions))]
92 self.tracer.lock_exclusive();
93
94 let mut inner = self.inner.lock();
95
96 if !inner.try_write() {
97 let event = InterruptibleEvent::new();
98 let guard = event.begin_wait();
99
100 inner.waiters.push_back(Waiter::Writer(event.clone()));
101
102 std::mem::drop(inner);
103
104 current_task.block_until(guard, zx::MonotonicInstant::INFINITE).map_err(|e| {
105 self.inner.lock().remove_waiter(&event);
106 e
107 })?;
108 }
109 Ok(RwQueueWriteGuard { queue: self })
110 }
111
112 #[cfg(any(test, debug_assertions))]
114 pub fn read_for_lock_ordering<'a>(&'a self) -> RwQueueReadGuard<'a, L> {
115 #[cfg(any(test, debug_assertions))]
118 self.tracer.lock_shared();
119
120 assert!(self.inner.lock().try_read(), "Cannot fail to acquire a read for lock ordering.");
121
122 RwQueueReadGuard { queue: self }
123 }
124}
125
126impl<L> Default for RwQueue<L> {
127 fn default() -> Self {
128 Self {
129 inner: Default::default(),
130 _phantom: PhantomData,
131 #[cfg(any(test, debug_assertions))]
132 tracer: Default::default(),
133 }
134 }
135}
136
137const READY: usize = 0;
139
140const WRITER: usize = 0b01;
142
143const READER: usize = 0b10;
145
146fn has_writer(state: usize) -> bool {
148 state & WRITER != 0
149}
150
151fn has_reader(state: usize) -> bool {
153 state >= READER
154}
155
156fn debug_assert_consistent(state: usize) {
157 debug_assert!(!has_writer(state) || !has_reader(state));
158}
159
160#[derive(Debug, Clone)]
161enum Waiter {
162 Reader(Arc<InterruptibleEvent>),
163 Writer(Arc<InterruptibleEvent>),
164}
165
166#[derive(Debug, Default)]
167struct RwQueueInner {
168 state: usize,
172
173 waiters: VecDeque<Waiter>,
175}
176
177impl RwQueueInner {
178 fn has_waiters(&self) -> bool {
179 !self.waiters.is_empty()
180 }
181
182 fn try_read(&mut self) -> bool {
183 debug_assert_consistent(self.state);
184 if !has_writer(self.state) && !self.has_waiters() {
185 if let Some(new_state) = self.state.checked_add(READER) {
186 self.state = new_state;
187 return true;
188 }
189 }
190 false
191 }
192
193 fn try_write(&mut self) -> bool {
194 debug_assert_consistent(self.state);
195 if self.state == READY && !self.has_waiters() {
196 self.state += WRITER;
197 true
198 } else {
199 false
200 }
201 }
202
203 fn unlock_read(&mut self) {
204 debug_assert!(has_reader(self.state) && !has_writer(self.state));
205 self.state -= READER;
206
207 if !has_reader(self.state) && self.has_waiters() {
208 self.notify_next();
209 }
210 }
211
212 fn unlock_write(&mut self) {
213 debug_assert!(has_writer(self.state) && !has_reader(self.state));
214 self.state -= WRITER;
215
216 if self.has_waiters() {
217 self.notify_next();
218 }
219 }
220
221 fn notify_next(&mut self) {
222 while let Some(waiter) = self.waiters.front() {
223 match waiter {
224 Waiter::Reader(reader) => {
225 if has_writer(self.state) {
226 return;
227 }
228 let Some(new_state) = self.state.checked_add(READER) else {
232 return;
233 };
234 self.state = new_state;
235 reader.notify();
236 }
237 Waiter::Writer(writer) => {
238 if has_reader(self.state) || has_writer(self.state) {
239 return;
240 }
241 self.state += WRITER;
244 writer.notify();
245 }
246 }
247 self.waiters.pop_front();
248 }
249 debug_assert_consistent(self.state);
250 }
251
252 fn remove_waiter(&mut self, event: &Arc<InterruptibleEvent>) {
253 self.waiters.retain(|waiter| {
254 let (Waiter::Reader(other) | Waiter::Writer(other)) = waiter;
255 !Arc::ptr_eq(event, other)
256 });
257 }
258}
259
260pub struct RwQueueReadGuard<'a, L> {
261 queue: &'a RwQueue<L>,
262}
263
264impl<'a, L> Drop for RwQueueReadGuard<'a, L> {
265 fn drop(&mut self) {
266 self.queue.unlock_read();
267 }
268}
269
270pub struct RwQueueWriteGuard<'a, L> {
271 queue: &'a RwQueue<L>,
272}
273
274impl<'a, L> Drop for RwQueueWriteGuard<'a, L> {
275 fn drop(&mut self) {
276 self.queue.unlock_write();
277 }
278}
279
280#[cfg(any(test, debug_assertions))]
281mod tracer {
282
283 #[derive(Debug, Default)]
284 pub struct FakeRwLock {}
285
286 #[allow(
287 clippy::undocumented_unsafe_blocks,
288 reason = "Force documented unsafe blocks in Starnix"
289 )]
290 unsafe impl lock_api::RawRwLock for FakeRwLock {
291 const INIT: Self = Self {};
292
293 type GuardMarker = lock_api::GuardNoSend;
294
295 fn lock_shared(&self) {}
296 fn try_lock_shared(&self) -> bool {
297 false
298 }
299 unsafe fn unlock_shared(&self) {}
300
301 fn lock_exclusive(&self) {}
302 fn try_lock_exclusive(&self) -> bool {
303 false
304 }
305 unsafe fn unlock_exclusive(&self) {}
306
307 fn is_locked(&self) -> bool {
308 false
309 }
310 }
311
312 pub type MutexTracer = tracing_mutex::lockapi::TracingWrapper<FakeRwLock>;
314}
315
316#[cfg(not(any(test, debug_assertions)))]
319use tracing_mutex as _;
320
321#[cfg(test)]
322mod test {
323 use super::*;
324 use crate::task::Kernel;
325 use crate::task::dynamic_thread_spawner::SpawnRequestBuilder;
326 use crate::testing::*;
327 use futures::executor::block_on;
328 use futures::future::join_all;
329 use starnix_sync::lock_ordering;
330 use std::future::Future;
331 use std::pin::Pin;
332 use std::sync::Barrier;
333 use std::sync::atomic::{AtomicUsize, Ordering};
334
335 #[::fuchsia::test]
336 fn test_remove_from_queue() {
337 let mut inner = RwQueueInner::default();
338 let event1 = InterruptibleEvent::new();
339 let event2 = InterruptibleEvent::new();
340 let event3 = InterruptibleEvent::new();
341 inner.waiters.push_back(Waiter::Writer(event1.clone()));
342 inner.waiters.push_back(Waiter::Writer(event2.clone()));
343 inner.waiters.push_back(Waiter::Writer(event3.clone()));
344
345 inner.remove_waiter(&event2);
346
347 let waiter = inner.waiters.pop_front().expect("should have a waiter");
348 let Waiter::Writer(event) = waiter else {
349 unreachable!();
350 };
351 assert!(Arc::ptr_eq(&event1, &event));
352
353 let waiter = inner.waiters.pop_front().expect("should have a waiter");
354 let Waiter::Writer(event) = waiter else {
355 unreachable!();
356 };
357 assert!(Arc::ptr_eq(&event3, &event));
358
359 assert!(inner.waiters.is_empty());
360 }
361
362 #[::fuchsia::test]
363 async fn test_write_and_read() {
364 lock_ordering! {
365 Unlocked => TestLevel
366 }
367
368 spawn_kernel_and_run(async |current_task| {
369 let queue = RwQueue::<TestLevel>::default();
370 let read_guard1 = queue.read(current_task).expect("shouldn't be interrupted");
371 std::mem::drop(read_guard1);
372
373 let write_guard = queue.write(current_task).expect("shouldn't be interrupted");
374 std::mem::drop(write_guard);
375
376 let read_guard2 = queue.read(current_task).expect("shouldn't be interrupted");
377 std::mem::drop(read_guard2);
378 })
379 .await;
380 }
381
382 #[::fuchsia::test]
383 async fn test_read_in_parallel() {
384 spawn_kernel_and_run(async |current_task| {
385 let kernel = current_task.kernel();
386 lock_ordering! {
387 Unlocked => TestLevel
388 }
389 struct Info {
390 barrier: Barrier,
391 queue: RwQueue<TestLevel>,
392 }
393
394 let info =
395 Arc::new(Info { barrier: Barrier::new(2), queue: RwQueue::<TestLevel>::default() });
396
397 let info1 = Arc::clone(&info);
398 let closure1 = move |current_task: &CurrentTask| {
399 let guard = info1.queue.read(current_task).expect("shouldn't be interrupted");
400 info1.barrier.wait();
401 std::mem::drop(guard);
402 };
403 let (thread1, req) =
404 SpawnRequestBuilder::new().with_sync_closure(closure1).build_with_async_result();
405 kernel.kthreads.spawner().spawn_from_request(req);
406
407 let info2 = Arc::clone(&info);
408 let closure2 = move |current_task: &CurrentTask| {
409 let guard = info2.queue.read(current_task).expect("shouldn't be interrupted");
410 info2.barrier.wait();
411 std::mem::drop(guard);
412 };
413 let (thread2, req) =
414 SpawnRequestBuilder::new().with_sync_closure(closure2).build_with_async_result();
415 kernel.kthreads.spawner().spawn_from_request(req);
416
417 block_on(async {
418 thread1.await.expect("failed to join thread");
419 thread2.await.expect("failed to join thread");
420 });
421 })
422 .await;
423 }
424
425 lock_ordering! {
426 Unlocked => A
427 }
428 struct State {
429 queue: RwQueue<A>,
430 gate: Barrier,
431 writer_count: AtomicUsize,
432 reader_count: AtomicUsize,
433 }
434
435 impl State {
436 fn new(n: usize) -> State {
437 State {
438 queue: Default::default(),
439 gate: Barrier::new(n),
440 writer_count: Default::default(),
441 reader_count: Default::default(),
442 }
443 }
444
445 fn spawn_writer(
446 state: Arc<Self>,
447 kernel: Arc<Kernel>,
448 count: usize,
449 ) -> Pin<Box<dyn Future<Output = Result<(), Errno>> + Send>> {
450 let closure = move |current_task: &CurrentTask| {
451 state.gate.wait();
452 for _ in 0..count {
453 let guard = state.queue.write(current_task).expect("shouldn't be interrupted");
454 let writer_count = state.writer_count.fetch_add(1, Ordering::Acquire) + 1;
455 let reader_count = state.reader_count.load(Ordering::Acquire);
456 state.writer_count.fetch_sub(1, Ordering::Release);
457 std::mem::drop(guard);
458 assert_eq!(writer_count, 1, "More than one writer held the lock at once.");
459 assert_eq!(
460 reader_count, 0,
461 "A reader and writer held the lock at the same time."
462 );
463 }
464 };
465 let (result, req) =
466 SpawnRequestBuilder::new().with_sync_closure(closure).build_with_async_result();
467 kernel.kthreads.spawner().spawn_from_request(req);
468 Box::pin(result)
469 }
470
471 fn spawn_reader(
472 state: Arc<Self>,
473 kernel: Arc<Kernel>,
474 count: usize,
475 ) -> Pin<Box<dyn Future<Output = Result<(), Errno>> + Send>> {
476 let closure = move |current_task: &CurrentTask| {
477 state.gate.wait();
478 for _ in 0..count {
479 let guard = state.queue.read(current_task).expect("shouldn't be interrupted");
480 let reader_count = state.reader_count.fetch_add(1, Ordering::Acquire) + 1;
481 let writer_count = state.writer_count.load(Ordering::Acquire);
482 state.reader_count.fetch_sub(1, Ordering::Release);
483 std::mem::drop(guard);
484 assert_eq!(
485 writer_count, 0,
486 "A reader and writer held the lock at the same time."
487 );
488 assert!(reader_count > 0, "A reader held the lock without being counted.");
489 }
490 };
491 let (result, req) =
492 SpawnRequestBuilder::new().with_sync_closure(closure).build_with_async_result();
493 kernel.kthreads.spawner().spawn_from_request(req);
494 Box::pin(result)
495 }
496 }
497
498 #[::fuchsia::test]
499 async fn test_thundering_reads_and_writes() {
500 spawn_kernel_and_run(async |current_task| {
501 let kernel = current_task.kernel();
502 const THREAD_PAIRS: usize = 10;
503
504 let state = Arc::new(State::new(THREAD_PAIRS * 2));
505 let mut threads = vec![];
506 for _ in 0..THREAD_PAIRS {
507 threads.push(State::spawn_writer(Arc::clone(&state), kernel.clone(), 100));
508 threads.push(State::spawn_reader(Arc::clone(&state), kernel.clone(), 100));
509 }
510
511 block_on(join_all(threads)).into_iter().for_each(|r| r.expect("failed to join thread"));
512 })
513 .await;
514 }
515}