1use std::{cell::Cell, fmt, marker::PhantomData, rc, task::Waker};
3
4#[derive(Default)]
21pub struct LocalWaker {
22 waker: Cell<Option<Waker>>,
23 _t: PhantomData<rc::Rc<()>>,
24}
25
26impl LocalWaker {
27 pub fn new() -> Self {
29 LocalWaker::with(None)
30 }
31
32 pub fn with(waker: Option<Waker>) -> Self {
34 LocalWaker {
35 waker: Cell::new(waker),
36 _t: PhantomData,
37 }
38 }
39
40 #[inline]
41 pub fn register(&self, waker: &Waker) -> bool {
45 match self.waker.take() {
46 Some(prev) if prev.will_wake(waker) => {
47 self.waker.set(Some(prev));
48 true
49 }
50 prev => {
51 self.waker.set(Some(waker.clone()));
52 prev.is_some()
53 }
54 }
55 }
56
57 #[inline]
58 pub fn wake(&self) {
63 if let Some(waker) = self.take() {
64 waker.wake();
65 }
66 }
67
68 #[inline]
69 pub fn wake_checked(&self) -> bool {
74 if let Some(waker) = self.take() {
75 waker.wake();
76 true
77 } else {
78 false
79 }
80 }
81
82 pub fn take(&self) -> Option<Waker> {
86 self.waker.take()
87 }
88
89 #[doc(hidden)]
90 pub fn is_set(&self) -> bool {
92 let waker = self.waker.take();
93 let set = waker.is_some();
94 self.waker.set(waker);
95 set
96 }
97}
98
99impl Clone for LocalWaker {
101 fn clone(&self) -> Self {
102 LocalWaker::new()
103 }
104}
105
106impl fmt::Debug for LocalWaker {
107 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
108 write!(f, "LocalWaker")
109 }
110}
111
112pub async fn yield_to() {
114 use std::{future::Future, pin::Pin, task::Context, task::Poll};
115
116 struct Yield {
117 completed: bool,
118 }
119
120 impl Future for Yield {
121 type Output = ();
122
123 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
124 if self.completed {
125 return Poll::Ready(());
126 }
127
128 self.completed = true;
129 cx.waker().wake_by_ref();
130
131 Poll::Pending
132 }
133 }
134
135 Yield { completed: false }.await;
136}
137
138#[cfg(test)]
139mod test {
140 use super::*;
141
142 #[ntex::test]
143 async fn yield_test() {
144 yield_to().await;
145 }
146}
147
148#[cfg(test)]
149mod tests {
150 use std::sync::atomic::{AtomicUsize, Ordering};
151 use std::task::{RawWaker, RawWakerVTable};
152
153 use super::*;
154
155 static CLONES: AtomicUsize = AtomicUsize::new(0);
156
157 static VTABLE: RawWakerVTable = RawWakerVTable::new(
158 |p| {
159 CLONES.fetch_add(1, Ordering::Relaxed);
160 RawWaker::new(p, &VTABLE)
161 },
162 |_| {},
163 |_| {},
164 |_| {},
165 );
166
167 #[test]
168 fn test_register_same_waker() {
169 static A: u8 = 0;
170 static B: u8 = 0;
171 let a = unsafe { Waker::from_raw(RawWaker::new((&raw const A).cast(), &VTABLE)) };
172 let b = unsafe { Waker::from_raw(RawWaker::new((&raw const B).cast(), &VTABLE)) };
173
174 let w = LocalWaker::new();
175 assert!(!w.register(&a));
176 assert!(w.register(&a));
177 assert!(w.register(&a));
178 assert_eq!(CLONES.load(Ordering::Relaxed), 1);
179
180 assert!(w.register(&b));
181 assert_eq!(CLONES.load(Ordering::Relaxed), 2);
182 assert!(w.take().unwrap().will_wake(&b));
183 }
184
185 #[test]
186 fn local_waker_clone_is_empty() {
187 let waker = LocalWaker::new();
188 waker.register(std::task::Waker::noop());
189 assert!(waker.is_set());
190 let cloned = waker.clone();
191 assert!(!cloned.is_set());
192 assert!(!cloned.wake_checked());
193 assert_eq!(format!("{waker:?}"), "LocalWaker");
194 }
195}