Skip to main content

ntex_util/services/
counter.rs

1use std::{cell::Cell, cell::RefCell, future::poll_fn, rc::Rc, task::Context, task::Poll};
2
3use crate::task::LocalWaker;
4
5/// A shared count with an asynchronous capacity notification.
6///
7/// Clones use the same count and capacity, but each clone registers its own
8/// waiting task.
9#[derive(Debug)]
10pub struct Counter(usize, Rc<CounterInner>);
11
12#[derive(Debug)]
13struct CounterInner {
14    count: Cell<usize>,
15    capacity: Cell<usize>,
16    tasks: RefCell<slab::Slab<LocalWaker>>,
17}
18
19impl Counter {
20    /// Creates a counter with the specified capacity.
21    pub fn new(capacity: usize) -> Self {
22        let mut tasks = slab::Slab::new();
23        let idx = tasks.insert(LocalWaker::new());
24
25        Counter(
26            idx,
27            Rc::new(CounterInner {
28                count: Cell::new(0),
29                capacity: Cell::new(capacity),
30                tasks: RefCell::new(tasks),
31            }),
32        )
33    }
34
35    /// Acquires one count and returns a guard that releases it on drop.
36    ///
37    /// This does not wait for capacity; call [`available`](Self::available)
38    /// first when exceeding the configured capacity is not acceptable.
39    pub fn get(&self) -> CounterGuard {
40        CounterGuard::new(self.1.clone())
41    }
42
43    /// Changes the capacity and wakes waiting tasks.
44    pub fn set_capacity(&self, cap: usize) {
45        self.1.capacity.set(cap);
46        self.1.notify();
47    }
48
49    /// Returns `true` if another count can be acquired without exceeding the capacity.
50    pub fn is_available(&self) -> bool {
51        self.1.count.get() < self.1.capacity.get()
52    }
53
54    /// Waits until the counter has free capacity.
55    ///
56    /// Returns immediately if there is capacity available. Otherwise,
57    /// registers the current task for wakeup and waits until a slot is freed.
58    pub async fn available(&self) {
59        poll_fn(|cx| {
60            if self.poll_available(cx) {
61                Poll::Ready(())
62            } else {
63                Poll::Pending
64            }
65        })
66        .await;
67    }
68
69    /// Waits until the counter reaches its capacity (i.e., becomes unavailable).
70    pub async fn unavailable(&self) {
71        poll_fn(|cx| {
72            if self.is_available() {
73                self.1.tasks.borrow()[self.0].register(cx.waker());
74                Poll::Pending
75            } else {
76                Poll::Ready(())
77            }
78        })
79        .await;
80    }
81
82    /// Check if counter is not at capacity. If counter at capacity
83    /// it registers notification for current task.
84    fn poll_available(&self, cx: &mut Context<'_>) -> bool {
85        if self.1.count.get() < self.1.capacity.get() {
86            true
87        } else {
88            let tasks = self.1.tasks.borrow();
89            tasks[self.0].register(cx.waker());
90            false
91        }
92    }
93
94    /// Returns the number of currently held guards.
95    pub fn total(&self) -> usize {
96        self.1.count.get()
97    }
98}
99
100impl Clone for Counter {
101    fn clone(&self) -> Self {
102        let idx = self.1.tasks.borrow_mut().insert(LocalWaker::new());
103        Self(idx, self.1.clone())
104    }
105}
106
107impl Drop for Counter {
108    fn drop(&mut self) {
109        self.1.tasks.borrow_mut().remove(self.0);
110    }
111}
112
113#[derive(Debug)]
114/// An acquired counter slot.
115///
116/// Dropping the guard releases the slot and wakes availability waiters.
117pub struct CounterGuard(Rc<CounterInner>);
118
119impl CounterGuard {
120    fn new(inner: Rc<CounterInner>) -> Self {
121        inner.inc();
122        CounterGuard(inner)
123    }
124}
125
126impl Unpin for CounterGuard {}
127
128impl Drop for CounterGuard {
129    fn drop(&mut self) {
130        self.0.dec();
131    }
132}
133
134impl CounterInner {
135    fn inc(&self) {
136        let num = self.count.get() + 1;
137        self.count.set(num);
138        if num == self.capacity.get() {
139            self.notify();
140        }
141    }
142
143    fn dec(&self) {
144        let num = self.count.get();
145        self.count.set(num - 1);
146        if num == self.capacity.get() {
147            self.notify();
148        }
149    }
150
151    fn notify(&self) {
152        let tasks = self.tasks.borrow();
153        for (_, task) in &*tasks {
154            task.wake();
155        }
156    }
157}
158
159#[cfg(test)]
160mod tests {
161    use std::time::Duration;
162
163    use super::*;
164    use crate::time::sleep;
165
166    #[ntex::test]
167    async fn test_unavailable_is_woken() {
168        let counter = Counter::new(2);
169        let done = Rc::new(Cell::new(false));
170
171        let (c, d) = (counter.clone(), done.clone());
172        crate::spawn(async move {
173            c.unavailable().await;
174            d.set(true);
175        });
176        sleep(Duration::from_millis(10)).await;
177        assert!(!done.get());
178
179        let _g1 = counter.get();
180        sleep(Duration::from_millis(10)).await;
181        assert!(!done.get());
182
183        let _g2 = counter.get();
184        sleep(Duration::from_millis(10)).await;
185        assert!(done.get());
186    }
187
188    #[ntex::test]
189    async fn test_available_is_woken() {
190        let counter = Counter::new(1);
191        let guard = counter.get();
192        let done = Rc::new(Cell::new(false));
193
194        let (c, d) = (counter.clone(), done.clone());
195        crate::spawn(async move {
196            c.available().await;
197            d.set(true);
198        });
199        sleep(Duration::from_millis(10)).await;
200        assert!(!done.get());
201
202        drop(guard);
203        sleep(Duration::from_millis(10)).await;
204        assert!(done.get());
205    }
206
207    #[ntex::test]
208    async fn test_set_capacity() {
209        let counter = Counter::new(1);
210        let guard = counter.get();
211        assert_eq!(counter.total(), 1);
212        assert!(!counter.is_available());
213
214        let counter2 = counter.clone();
215        let hnd = crate::spawn(async move { counter2.available().await });
216        sleep(Duration::from_millis(10)).await;
217        assert!(!hnd.is_finished());
218
219        // raising the capacity wakes waiters
220        counter.set_capacity(2);
221        assert!(counter.is_available());
222        crate::time::timeout(Duration::from_secs(1), hnd)
223            .await
224            .unwrap()
225            .unwrap();
226
227        drop(guard);
228        assert_eq!(counter.total(), 0);
229    }
230}