Skip to main content

starnix_core/vfs/
rw_queue.rs

1// Copyright 2023 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use 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    // Used to inform our deadlock detector about the waiters in the queue.
25    #[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        // TODO(https://fxbug.dev/532501001): A lock dep guard should
63        // be added once FsNodeAppend is correct in respect of LockDep.
64        #[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        // TODO(https://fxbug.dev/532501001): A lock dep guard should
90        // be added once FsNodeAppend is correct in respect of LockDep.
91        #[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    /// Used to establish lock ordering.
113    #[cfg(any(test, debug_assertions))]
114    pub fn read_for_lock_ordering<'a>(&'a self) -> RwQueueReadGuard<'a, L> {
115        // TODO(https://fxbug.dev/532501001): A lock dep guard should
116        // be added once FsNodeAppend is correct in respect of LockDep.
117        #[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
137/// The queue is ready for any operation.
138const READY: usize = 0;
139
140/// The queue has exactly one writer.
141const WRITER: usize = 0b01;
142
143/// Each writer in the queue increments the state by this amount.
144const READER: usize = 0b10;
145
146/// A writer is currently running.
147fn has_writer(state: usize) -> bool {
148    state & WRITER != 0
149}
150
151/// At elast one reader is currently running.
152fn 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    /// What operations are currently ongoing.
169    ///
170    /// See READY, READER, WRITER above for what these bits mean.
171    state: usize,
172
173    /// The operations that are waiting for the ongoing operations to complete.
174    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                    // We need to use `checked_add` to ensure we do not
229                    // overflow the number of readers. If that happens, we just
230                    // need to wait for the enormous number of readers to finish.
231                    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                    // We can never overflow writers because we only let one
242                    // through at a time.
243                    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    // We should replace this type with tracing_mutex::MutexId once that type is public.
313    pub type MutexTracer = tracing_mutex::lockapi::TracingWrapper<FakeRwLock>;
314}
315
316// We use tracing_mutex in tests and debug assertions, but we don't want to pull it in for
317// production.
318#[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}