Skip to main content

conc_util/event/
onceevent.rs

1use core::{
2    mem::ManuallyDrop,
3    ops::Deref,
4    pin::Pin,
5    task::{self, Waker},
6};
7
8use alloc::vec::Vec;
9
10use crate::{
11    event::Event,
12    sync::{OnceFlag, SpinLock},
13};
14
15#[derive(Debug)]
16pub struct OnceEvent {
17    used: OnceFlag,
18    wakers: SpinLock<ManuallyDrop<Vec<Waker>>>,
19}
20
21impl OnceEvent {
22    pub fn new() -> Self {
23        Self {
24            used: OnceFlag::new(),
25            wakers: SpinLock::new(ManuallyDrop::new(Vec::new())),
26        }
27    }
28}
29
30impl<D: Deref<Target = OnceEvent>> Event for D {
31    fn fire(self) {
32        if self.used.fire() {
33            return;
34        }
35
36        let mut lock = self.wakers.lock();
37        // SAFETY:
38        // 1) Value is taken once - `used` is checked at the beginning
39        // 2) Value is never used again - `poll` checks the same `used`
40        let wakers = unsafe { ManuallyDrop::take(&mut lock) };
41        drop(lock);
42
43        for waker in wakers {
44            waker.wake();
45        }
46    }
47
48    fn wait(self) -> impl Future<Output = ()> {
49        OnceEventWaiter(self)
50    }
51}
52
53struct OnceEventWaiter<D>(D);
54
55impl<D: Deref<Target = OnceEvent>> Future for OnceEventWaiter<D> {
56    type Output = ();
57
58    fn poll(self: Pin<&mut Self>, cx: &mut task::Context<'_>) -> task::Poll<Self::Output> {
59        if self.0.used.fired() {
60            return task::Poll::Ready(());
61        }
62
63        let mut wakers = self.0.wakers.lock();
64
65        if self.0.used.fired() {
66            return task::Poll::Ready(());
67        }
68
69        wakers.push(cx.waker().clone());
70        return task::Poll::Pending;
71    }
72}
73
74impl Drop for OnceEvent {
75    fn drop(&mut self) {
76        if !self.used.fired_mut() {
77            // SAFETY: this is only dropped here and taken in `fire()`,
78            // but it also sets the `used` flag which is checked above.
79            unsafe { ManuallyDrop::drop(self.wakers.get()) };
80        }
81    }
82}