Skip to main content

ntex_util/channel/
condition.rs

1use std::{cell, fmt, future::Future, future::poll_fn, pin::Pin, task::Context, task::Poll};
2
3use slab::Slab;
4
5use super::cell::Cell;
6use crate::task::LocalWaker;
7
8#[derive(Clone, Debug, PartialEq, Eq)]
9/// Result produced by a [`Condition`] waiter.
10pub enum ConditionResult<T> {
11    /// The condition delivered a value.
12    Value(T),
13    /// The condition has been locked and will not deliver more values.
14    Locked,
15    /// The last handle to the condition was dropped.
16    Dropped,
17}
18
19#[derive(Copy, Clone, PartialEq, Eq, Debug)]
20enum State {
21    Normal,
22    Locked,
23    Dropped,
24}
25
26/// A condition that can wake several waiting tasks at once.
27///
28/// Notifications are not queued. A waiter must have been polled and registered
29/// its waker before [`notify`](Self::notify) is called, otherwise it misses that
30/// value. Use [`notify_and_lock`](Self::notify_and_lock) when no later
31/// notifications should be accepted.
32pub struct Condition<T = ()> {
33    inner: Cell<Inner<T>>,
34}
35
36/// A task waiting for a [`Condition`] notification.
37pub struct Waiter<T = ()> {
38    token: usize,
39    inner: Cell<Inner<T>>,
40}
41
42struct Inner<T> {
43    data: Slab<Option<Item<T>>>,
44    count: usize,
45    state: State,
46}
47
48struct Item<T> {
49    val: cell::Cell<ConditionResult<T>>,
50    waker: LocalWaker,
51}
52
53impl Default for Condition<()> {
54    fn default() -> Self {
55        Self::new()
56    }
57}
58
59impl<T> Clone for Condition<T> {
60    fn clone(&self) -> Self {
61        let inner = self.inner.clone();
62        inner.get_mut().count += 1;
63        Self { inner }
64    }
65}
66
67impl<T> fmt::Debug for Condition<T> {
68    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
69        f.debug_struct("Condition")
70            .field("state", &self.inner.get_ref().state)
71            .finish()
72    }
73}
74
75impl<T> Condition<T> {
76    /// Creates an unlocked condition with no waiters.
77    pub fn new() -> Condition<T> {
78        Condition {
79            inner: Cell::new(Inner {
80                data: Slab::new(),
81                count: 1,
82                state: State::Normal,
83            }),
84        }
85    }
86}
87
88impl<T: Clone> Condition<T> {
89    /// Creates a new waiter.
90    ///
91    /// The waiter starts listening when it is first polled, not when this
92    /// method returns.
93    pub fn wait(&self) -> Waiter<T> {
94        let token = self.inner.get_mut().data.insert(None);
95        Waiter {
96            token,
97            inner: self.inner.clone(),
98        }
99    }
100
101    /// Sends `val` to every waiter that is currently being polled.
102    ///
103    /// The value is cloned for each registered waiter. Unpolled waiters do not
104    /// receive it, and the value is not retained for future waiters.
105    pub fn notify(&self, val: T) {
106        let inner = self.inner.get_ref();
107        if inner.state != State::Normal {
108            return;
109        }
110        for (_, item) in &inner.data {
111            if let Some(item) = item
112                && item.waker.wake_checked()
113            {
114                item.val.set(ConditionResult::Value(val.clone()));
115            }
116        }
117    }
118
119    /// Notifies the current waiters and permanently locks the condition.
120    ///
121    /// Registered waiters receive `val`. Later readiness checks return
122    /// [`ConditionResult::Locked`], and later calls to [`notify`](Self::notify)
123    /// do not deliver another value.
124    pub fn notify_and_lock(&self, val: T) {
125        self.notify(val);
126        self.inner.get_mut().state = State::Locked;
127    }
128}
129
130impl<T: Default> Condition<T> {
131    /// Sends `T::default()` to every waiter that is currently being polled.
132    pub fn notify_default(&self) {
133        let inner = self.inner.get_ref();
134        if inner.state != State::Normal {
135            return;
136        }
137        for (_, item) in &inner.data {
138            if let Some(item) = item
139                && item.waker.wake_checked()
140            {
141                item.val.set(ConditionResult::Value(T::default()));
142            }
143        }
144    }
145}
146
147impl<T> Drop for Condition<T> {
148    fn drop(&mut self) {
149        let inner = self.inner.get_mut();
150        inner.count -= 1;
151        if inner.count == 0 {
152            inner.state = State::Dropped;
153            for (_, item) in &inner.data {
154                if let Some(item) = item
155                    && item.waker.wake_checked()
156                {
157                    item.val.set(ConditionResult::Dropped);
158                }
159            }
160        }
161    }
162}
163
164impl<T> Waiter<T> {
165    /// Waits for the next condition result.
166    pub async fn ready(&self) -> ConditionResult<T> {
167        poll_fn(|cx| self.poll_ready(cx)).await
168    }
169
170    /// Polls for the next condition result.
171    ///
172    /// The first poll registers this waiter. While the condition remains
173    /// unlocked, later notifications wake the registered task.
174    pub fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<ConditionResult<T>> {
175        let parent = self.inner.get_mut();
176        let inner = unsafe { parent.data.get_unchecked_mut(self.token) };
177
178        if inner.is_none() {
179            if parent.state == State::Normal {
180                let waker = LocalWaker::default();
181                waker.register(cx.waker());
182                *inner = Some(Item {
183                    waker,
184                    val: cell::Cell::new(ConditionResult::Locked),
185                });
186                return Poll::Pending;
187            }
188        } else {
189            let item = inner.as_mut().unwrap();
190            if !item.waker.register(cx.waker()) {
191                return Poll::Ready(item.val.replace(ConditionResult::Locked));
192            }
193        }
194
195        match parent.state {
196            State::Normal => Poll::Pending,
197            State::Locked => Poll::Ready(ConditionResult::Locked),
198            State::Dropped => Poll::Ready(ConditionResult::Dropped),
199        }
200    }
201}
202
203impl<T> Clone for Waiter<T> {
204    fn clone(&self) -> Self {
205        let token = self.inner.get_mut().data.insert(None);
206        Waiter {
207            token,
208            inner: self.inner.clone(),
209        }
210    }
211}
212
213impl<T> Future for Waiter<T> {
214    type Output = ConditionResult<T>;
215
216    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
217        self.get_mut().poll_ready(cx)
218    }
219}
220
221impl<T> fmt::Debug for Waiter<T> {
222    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
223        f.debug_struct("Waiter").finish()
224    }
225}
226
227impl<T> Drop for Waiter<T> {
228    fn drop(&mut self) {
229        self.inner.get_mut().data.remove(self.token);
230    }
231}
232
233#[cfg(test)]
234mod tests {
235    use super::*;
236    use crate::future::lazy;
237
238    #[ntex::test]
239    #[allow(clippy::unit_cmp)]
240    async fn test_condition() {
241        let cond = Condition::<()>::new();
242        let mut waiter = cond.wait();
243        assert_eq!(
244            lazy(|cx| Pin::new(&mut waiter).poll(cx)).await,
245            Poll::Pending
246        );
247        cond.notify_default();
248        assert!(format!("{cond:?}").contains("Condition"));
249        assert!(format!("{waiter:?}").contains("Waiter"));
250        assert_eq!(waiter.await, ConditionResult::Value(()));
251
252        let mut waiter = cond.wait();
253        assert_eq!(
254            lazy(|cx| Pin::new(&mut waiter).poll(cx)).await,
255            Poll::Pending
256        );
257        let mut waiter2 = waiter.clone();
258        assert_eq!(
259            lazy(|cx| Pin::new(&mut waiter2).poll(cx)).await,
260            Poll::Pending
261        );
262
263        drop(cond);
264        assert_eq!(waiter.await, ConditionResult::Dropped);
265        assert_eq!(waiter2.await, ConditionResult::Dropped);
266    }
267
268    #[ntex::test]
269    async fn test_condition_poll() {
270        let cond = Condition::default().clone();
271        let waiter = cond.wait();
272        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
273        cond.notify_default();
274        waiter.ready().await;
275
276        let waiter2 = waiter.clone();
277        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
278        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
279        assert_eq!(lazy(|cx| waiter2.poll_ready(cx)).await, Poll::Pending);
280        assert_eq!(lazy(|cx| waiter2.poll_ready(cx)).await, Poll::Pending);
281
282        drop(cond);
283        assert_eq!(
284            lazy(|cx| waiter.poll_ready(cx)).await,
285            Poll::Ready(ConditionResult::Dropped)
286        );
287        assert_eq!(
288            lazy(|cx| waiter.poll_ready(cx)).await,
289            Poll::Ready(ConditionResult::Dropped)
290        );
291        assert_eq!(
292            lazy(|cx| waiter2.poll_ready(cx)).await,
293            Poll::Ready(ConditionResult::Dropped)
294        );
295        assert_eq!(
296            lazy(|cx| waiter2.poll_ready(cx)).await,
297            Poll::Ready(ConditionResult::Dropped)
298        );
299    }
300
301    #[ntex::test]
302    async fn test_condition_with() {
303        let cond = Condition::<String>::new();
304        let waiter = cond.wait();
305        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
306        cond.notify("TEST".into());
307        assert_eq!(
308            waiter.ready().await,
309            ConditionResult::Value("TEST".to_string())
310        );
311
312        let waiter2 = waiter.clone();
313        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
314        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
315        assert_eq!(lazy(|cx| waiter2.poll_ready(cx)).await, Poll::Pending);
316        assert_eq!(lazy(|cx| waiter2.poll_ready(cx)).await, Poll::Pending);
317
318        drop(cond);
319        assert_eq!(
320            lazy(|cx| waiter.poll_ready(cx)).await,
321            Poll::Ready(ConditionResult::Dropped)
322        );
323        assert_eq!(
324            lazy(|cx| waiter.poll_ready(cx)).await,
325            Poll::Ready(ConditionResult::Dropped)
326        );
327        assert_eq!(
328            lazy(|cx| waiter2.poll_ready(cx)).await,
329            Poll::Ready(ConditionResult::Dropped)
330        );
331        assert_eq!(
332            lazy(|cx| waiter2.poll_ready(cx)).await,
333            Poll::Ready(ConditionResult::Dropped)
334        );
335    }
336
337    #[ntex::test]
338    async fn waiter_future_does_not_require_default() {
339        #[derive(Clone, Debug, PartialEq, Eq)]
340        struct Value(&'static str);
341
342        let cond = Condition::<Value>::new();
343        let waiter = cond.wait();
344        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
345
346        cond.notify(Value("ready"));
347        assert_eq!(waiter.await, ConditionResult::Value(Value("ready")));
348    }
349
350    #[ntex::test]
351    async fn notify_ready() {
352        let cond = Condition::default().clone();
353        let waiter = cond.wait();
354        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
355
356        cond.notify_and_lock(());
357        assert_eq!(
358            lazy(|cx| waiter.poll_ready(cx)).await,
359            Poll::Ready(ConditionResult::Value(()))
360        );
361        assert_eq!(
362            lazy(|cx| waiter.poll_ready(cx)).await,
363            Poll::Ready(ConditionResult::Locked)
364        );
365        assert_eq!(
366            lazy(|cx| waiter.poll_ready(cx)).await,
367            Poll::Ready(ConditionResult::Locked)
368        );
369        cond.notify(());
370        assert_eq!(
371            lazy(|cx| waiter.poll_ready(cx)).await,
372            Poll::Ready(ConditionResult::Locked)
373        );
374
375        let waiter2 = cond.wait();
376        assert_eq!(
377            lazy(|cx| waiter2.poll_ready(cx)).await,
378            Poll::Ready(ConditionResult::Locked)
379        );
380    }
381
382    #[ntex::test]
383    async fn notify_with_and_lock_ready() {
384        // with
385        let cond = Condition::<String>::new();
386        let waiter = cond.wait();
387        let waiter2 = cond.wait();
388        assert_eq!(lazy(|cx| waiter.poll_ready(cx)).await, Poll::Pending);
389        assert_eq!(lazy(|cx| waiter2.poll_ready(cx)).await, Poll::Pending);
390
391        cond.notify_and_lock("TEST".into());
392        assert_eq!(
393            lazy(|cx| waiter.poll_ready(cx)).await,
394            Poll::Ready(ConditionResult::Value("TEST".into()))
395        );
396        assert_eq!(
397            lazy(|cx| waiter.poll_ready(cx)).await,
398            Poll::Ready(ConditionResult::Locked)
399        );
400        assert_eq!(
401            lazy(|cx| waiter.poll_ready(cx)).await,
402            Poll::Ready(ConditionResult::Locked)
403        );
404        assert_eq!(
405            lazy(|cx| waiter2.poll_ready(cx)).await,
406            Poll::Ready(ConditionResult::Value("TEST".into()))
407        );
408        assert_eq!(
409            lazy(|cx| waiter2.poll_ready(cx)).await,
410            Poll::Ready(ConditionResult::Locked)
411        );
412
413        let waiter2 = cond.wait();
414        assert_eq!(
415            lazy(|cx| waiter2.poll_ready(cx)).await,
416            Poll::Ready(ConditionResult::Locked)
417        );
418    }
419
420    #[ntex::test]
421    async fn notify_default_locked() {
422        let cond = Condition::<()>::new();
423        cond.notify_and_lock(());
424        let mut waiter = cond.wait();
425        cond.notify_default();
426        assert_eq!(
427            lazy(|cx| Pin::new(&mut waiter).poll(cx)).await,
428            Poll::Ready(ConditionResult::Locked)
429        );
430    }
431}