1use std::{borrow::Cow, cell::Cell, pin::Pin, task::Context, task::Poll, task::Waker};
3
4use ntex_util::HashMap;
5
6use crate::{IoRef, utils::Extensions};
7
8const NONE: u16 = u16::MAX;
9
10pub(crate) const TAG_DISCONNECT: usize = usize::MAX - 1;
12pub(crate) const TAG_WRITE: usize = usize::MAX - 2;
14
15pub(crate) fn check_public_tag(tag: usize) {
17 debug_assert!(tag < TAG_WRITE, "waker tag {tag} is reserved");
18}
19
20#[derive(Copy, Clone, Debug, PartialEq, Eq)]
25pub(crate) struct WaiterId {
26 idx: u16,
27 generation: u16,
28}
29
30#[derive(Debug)]
34pub(crate) struct WaiterEntry {
35 pub(crate) tag: usize,
36 pub(crate) id: Cell<Option<WaiterId>>,
37}
38
39impl WaiterEntry {
40 pub(crate) const fn new(tag: usize) -> Self {
41 Self {
42 tag,
43 id: Cell::new(None),
44 }
45 }
46
47 fn take(&self) -> Self {
48 Self {
49 tag: self.tag,
50 id: Cell::new(self.id.take()),
51 }
52 }
53}
54
55struct Entry {
56 waker: Option<Waker>,
57 prev: u16,
58 next: u16,
59 generation: u16,
60}
61
62pub(crate) struct Waiters {
69 free: u16,
70 entries: Vec<Entry>,
71 tags: HashMap<usize, u16>,
73}
74
75impl Default for Waiters {
76 fn default() -> Self {
77 Self::new()
78 }
79}
80
81impl Waiters {
82 pub(crate) fn new() -> Self {
83 Self {
84 free: NONE,
85 entries: Vec::new(),
86 tags: HashMap::default(),
87 }
88 }
89
90 pub(crate) fn register(&mut self, tag: usize, waker: &Waker) -> WaiterId {
92 let head = self.tags.get(&tag).copied().unwrap_or(NONE);
93
94 let idx = if self.free == NONE {
95 let idx = u16::try_from(self.entries.len()).expect("too many wakers");
96 assert!(idx != NONE, "too many wakers");
97 self.entries.push(Entry {
98 waker: Some(waker.clone()),
99 generation: 0,
100 prev: NONE,
101 next: head,
102 });
103 idx
104 } else {
105 let idx = self.free;
106 let entry = &mut self.entries[idx as usize];
107 self.free = entry.next;
108 entry.waker = Some(waker.clone());
109 entry.prev = NONE;
110 entry.next = head;
111 idx
112 };
113
114 if head != NONE {
115 self.entries[head as usize].prev = idx;
116 }
117 self.tags.insert(tag, idx);
118 WaiterId {
119 idx,
120 generation: self.entries[idx as usize].generation,
121 }
122 }
123
124 pub(crate) fn update(&mut self, id: WaiterId, waker: &Waker) -> bool {
128 if let Some(entry) = self.get(id)
129 && let Some(ref mut w) = entry.waker
130 {
131 w.clone_from(waker);
132 true
133 } else {
134 false
135 }
136 }
137
138 pub(crate) fn remove(&mut self, id: WaiterId, tag: usize) {
143 let Some(entry) = self.get(id) else {
144 return;
145 };
146 let (prev, next) = (entry.prev, entry.next);
147 if prev != NONE {
148 self.entries[prev as usize].next = next;
149 } else if next == NONE {
150 self.tags.remove(&tag);
151 } else {
152 self.tags.insert(tag, next);
153 }
154 if next != NONE {
155 self.entries[next as usize].prev = prev;
156 }
157 drop(self.release(id.idx));
158 }
159
160 pub(crate) fn wake_all(&mut self) {
162 while let Some(&tag) = self.tags.keys().next() {
163 self.wake(tag);
164 }
165 }
166
167 pub(crate) fn wake(&mut self, tag: usize) {
169 let Some(mut idx) = self.tags.remove(&tag) else {
170 return;
171 };
172
173 while idx != NONE {
174 let next = self.entries[idx as usize].next;
175 if let Some(waker) = self.release(idx) {
176 waker.wake();
177 }
178 idx = next;
179 }
180 }
181
182 #[cfg(test)]
183 pub(crate) fn len(&self) -> usize {
184 self.entries.iter().filter(|e| e.waker.is_some()).count()
185 }
186
187 #[cfg(test)]
188 fn is_registered(&self, id: WaiterId) -> bool {
189 self.entries
190 .get(id.idx as usize)
191 .is_some_and(|e| e.generation == id.generation && e.waker.is_some())
192 }
193
194 fn get(&mut self, id: WaiterId) -> Option<&mut Entry> {
195 self.entries
196 .get_mut(id.idx as usize)
197 .filter(|e| e.generation == id.generation && e.waker.is_some())
198 }
199
200 fn release(&mut self, idx: u16) -> Option<Waker> {
202 let entry = &mut self.entries[idx as usize];
203 entry.generation = entry.generation.wrapping_add(1);
204 entry.prev = NONE;
205 entry.next = self.free;
206 self.free = idx;
207 entry.waker.take()
208 }
209}
210
211pub(crate) struct WriteGuard<'a> {
213 ext: &'a Extensions,
214 slot: WaiterEntry,
215}
216
217impl<'a> WriteGuard<'a> {
218 pub(crate) fn new(ext: &'a Extensions) -> Self {
219 Self {
220 ext,
221 slot: WaiterEntry::new(TAG_WRITE),
222 }
223 }
224
225 pub(crate) fn register(&self, cx: &mut Context<'_>) {
226 self.ext.register_waker(&self.slot, cx.waker());
227 }
228}
229
230impl Drop for WriteGuard<'_> {
231 fn drop(&mut self) {
232 self.ext.remove_waker(&self.slot);
233 }
234}
235
236#[derive(Debug)]
246#[must_use = "a waiter does nothing unless polled"]
247pub struct Waiter<'a> {
248 io: Cow<'a, IoRef>,
249 waiter: WaiterEntry,
250}
251
252impl<'a> Waiter<'a> {
253 pub fn new(io: &'a IoRef, tag: usize) -> Self {
260 check_public_tag(tag);
261 Self {
262 io: Cow::Borrowed(io),
263 waiter: WaiterEntry::new(tag),
264 }
265 }
266
267 pub(crate) fn new_static(io: IoRef, tag: usize) -> Self {
268 Self {
269 io: Cow::Owned(io),
270 waiter: WaiterEntry::new(tag),
271 }
272 }
273
274 pub fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<()> {
278 let st = &self.io.0;
279 if st.flags.is_closed() {
280 Poll::Ready(())
281 } else {
282 st.extensions.poll_waker(&self.waiter, cx.waker())
283 }
284 }
285
286 pub fn into_static(self) -> Waiter<'static> {
290 let io = Cow::Owned(IoRef::clone(&self.io));
291
292 Waiter {
293 io,
294 waiter: self.waiter.take(),
295 }
296 }
297}
298
299impl Clone for Waiter<'_> {
300 fn clone(&self) -> Self {
302 Self {
303 io: self.io.clone(),
304 waiter: WaiterEntry::new(self.waiter.tag),
305 }
306 }
307}
308
309impl Future for Waiter<'_> {
310 type Output = ();
311
312 #[inline]
313 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
314 self.poll_ready(cx)
315 }
316}
317
318impl Drop for Waiter<'_> {
319 fn drop(&mut self) {
320 self.io.0.extensions.remove_waker(&self.waiter);
321 }
322}
323
324#[cfg(test)]
325mod tests {
326 use std::sync::{Arc, atomic::AtomicUsize, atomic::Ordering};
327 use std::task::{Wake, Waker};
328
329 use super::*;
330
331 struct Counter(AtomicUsize);
332
333 impl Wake for Counter {
334 fn wake(self: Arc<Self>) {
335 self.0.fetch_add(1, Ordering::Relaxed);
336 }
337 }
338
339 fn waker() -> (Arc<Counter>, Waker) {
340 let cnt = Arc::new(Counter(AtomicUsize::new(0)));
341 (cnt.clone(), Waker::from(cnt))
342 }
343
344 fn count(cnt: &Counter) -> usize {
345 cnt.0.load(Ordering::Relaxed)
346 }
347
348 #[test]
349 fn wake_by_tag() {
350 let mut wakers = Waiters::new();
351 let (c1, w1) = waker();
352 let (c2, w2) = waker();
353 let (c3, w3) = waker();
354
355 let id1 = wakers.register(0, &w1);
356 let id2 = wakers.register(1, &w2);
357 let id3 = wakers.register(0, &w3);
358
359 wakers.wake(0);
360 assert_eq!((count(&c1), count(&c2), count(&c3)), (1, 0, 1));
361 assert!(!wakers.is_registered(id1));
362 assert!(wakers.is_registered(id2));
363 assert!(!wakers.is_registered(id3));
364
365 wakers.wake(0);
367 assert_eq!((count(&c1), count(&c3)), (1, 1));
368
369 wakers.wake(1);
370 assert_eq!(count(&c2), 1);
371 assert!(!wakers.is_registered(id2));
372 assert_eq!(wakers.entries.len(), 3);
373 assert_eq!(wakers.len(), 0);
374 }
375
376 #[test]
377 fn dynamic_tags() {
378 let mut wakers = Waiters::default();
379 let (cnt, w) = waker();
380
381 wakers.wake(7);
383 assert!(wakers.tags.is_empty());
384
385 let id = wakers.register(200, &w);
386 assert_eq!(wakers.tags.len(), 1);
387 let id3 = wakers.register(3, &w);
388 wakers.register(3, &w);
389 assert_eq!(wakers.tags.len(), 2);
390
391 wakers.wake(3);
392 assert_eq!(count(&cnt), 2);
393 assert!(wakers.is_registered(id));
394 assert_eq!(wakers.tags.len(), 1);
395
396 wakers.remove(id3, 3);
398 wakers.remove(id, 200);
399 assert!(wakers.tags.is_empty());
400 wakers.wake(200);
401 assert_eq!(count(&cnt), 2);
402
403 wakers.register(3, &w);
404 let id3 = wakers.register(3, &w);
405 wakers.remove(id3, 3);
406 assert_eq!(wakers.tags.len(), 1);
407 wakers.wake(3);
408 assert_eq!(count(&cnt), 3);
409 assert!(wakers.tags.is_empty());
410 }
411
412 #[test]
413 fn capacity() {
414 let mut wakers = Waiters::new();
415 let (_, w) = waker();
416
417 let ids: Vec<_> = (0..u16::MAX).map(|_| wakers.register(0, &w)).collect();
418 assert_eq!(wakers.len(), usize::from(u16::MAX));
419
420 wakers.remove(ids[100], 0);
422 assert_eq!(wakers.register(0, &w).idx, 100);
423
424 let res = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
425 wakers.register(0, &w);
426 }));
427 assert!(res.is_err());
428 }
429
430 #[test]
431 fn remove() {
432 let mut wakers = Waiters::new();
433 let (cnt, w) = waker();
434
435 let ids: Vec<_> = (0..5).map(|_| wakers.register(0, &w)).collect();
437 wakers.remove(ids[4], 0);
438 wakers.remove(ids[2], 0);
439 wakers.remove(ids[0], 0);
440 assert_eq!(wakers.len(), 2);
441
442 wakers.remove(ids[0], 0);
444 assert_eq!(wakers.len(), 2);
445
446 wakers.wake(0);
447 assert_eq!(count(&cnt), 2);
448 assert_eq!(wakers.len(), 0);
449 }
450
451 #[test]
452 fn update() {
453 let mut wakers = Waiters::new();
454 let (c1, w1) = waker();
455 let (c2, w2) = waker();
456
457 let id = wakers.register(0, &w1);
458 assert!(wakers.update(id, &w2));
459 wakers.wake(0);
460 assert_eq!((count(&c1), count(&c2)), (0, 1));
461 assert!(!wakers.update(id, &w1));
462 }
463
464 #[test]
465 fn wake_all() {
466 let mut wakers = Waiters::new();
467 let (cnt, w) = waker();
468
469 wakers.wake_all();
470 for tag in [0, 1, 5, 5] {
471 wakers.register(tag, &w);
472 }
473 wakers.wake_all();
474 assert_eq!(count(&cnt), 4);
475 assert!(wakers.tags.is_empty());
476 assert_eq!(wakers.len(), 0);
477 }
478
479 #[test]
480 fn stale_ids() {
481 let mut wakers = Waiters::new();
482 let (c1, w1) = waker();
483 let (c2, w2) = waker();
484
485 let stale = wakers.register(0, &w1);
486 wakers.wake(0);
487 assert_eq!(count(&c1), 1);
488
489 let id = wakers.register(0, &w2);
491 assert_eq!(id.idx, stale.idx);
492 assert!(!wakers.is_registered(stale));
493 assert!(!wakers.update(stale, &w1));
494 wakers.remove(stale, 0);
495 assert!(wakers.is_registered(id));
496
497 wakers.wake(0);
498 assert_eq!((count(&c1), count(&c2)), (1, 1));
499 }
500
501 #[test]
502 fn free_list_reuse() {
503 let mut wakers = Waiters::new();
504 let (cnt, w) = waker();
505
506 let a: Vec<_> = (0..4).map(|i| wakers.register(i % 2, &w)).collect();
507 wakers.wake(0);
508 wakers.remove(a[1], 1);
509
510 let b: Vec<_> = (0..3).map(|_| wakers.register(1, &w)).collect();
512 assert_eq!(wakers.entries.len(), 4);
513 for id in &b {
514 assert!(wakers.is_registered(*id));
515 }
516 assert!(wakers.is_registered(a[3]));
517
518 wakers.wake(1);
519 assert_eq!(count(&cnt), 2 + 4);
520 assert_eq!(wakers.len(), 0);
521 }
522}