Skip to main content

tor_async_utils/bw_pool/
refiller.rs

1//! The refiller code that is the [`BandwidthRefiller`] is a public entity that needs to
2//! run into its own task and refill the associated [`super::BandwidthPool`] and serve
3//! any `RefillWaiter` pending on the pool to be replenished.
4//!
5//! # Driving a pool
6//!
7//! The refiller must live in its own task that owns the clock and the rate. Simply call
8//! [`BandwidthRefiller::run`] into a spawned task. It builds a token bucket from the given
9//! rate + burst and a [`tor_rtcompat::SleepProvider`].
10//!
11//! # Example
12//!
13//! ```ignore
14//! let (pool, refiller) = BandwidthPool::new(capacity);
15//! let config = TokenBucketConfig { rate, bucket_max: capacity };
16//! runtime.spawn(refiller.run(runtime.clone(), config))?;
17//! ```
18
19use futures::StreamExt as _;
20use futures::channel::mpsc;
21use futures::task::AtomicWaker;
22use std::num::NonZero;
23use std::sync::Arc;
24use std::sync::atomic::{AtomicU64, Ordering};
25use std::task::Waker;
26
27use tor_basic_utils::token_bucket::{TokenBucket, TokenBucketConfig};
28use tor_rtcompat::SleepProvider;
29
30use super::bucket::AtomicTokenBucket;
31
32/// A bandwidth request sent on the queue to the [`BandwidthRefiller`].
33///
34/// The amount of tokens wanted is fixed for the lifetime of a request and is only read
35/// by the refiller and thus why is travels along the waiter.
36pub(super) type BwRequest = (u64, Arc<RefillWaiter>);
37
38/// A refill waiter object through which the [`BandwidthRefiller`] signals a blocked
39/// acquirer.
40///
41/// A [`std::task::Waker`] doesn't carry any information back to the task so this waiter
42/// carries the number of granted tokens which doubles as the "was I granted" flag.
43///
44/// This lives in a [`super::BandwidthAcquirer`] and is reused at every acquire which
45/// means that once the task is launched, the steady state has no extra allocation.
46///
47/// For sake of simplicity, there is no cancellation path and so any granted bandwidth
48/// before cancellation (drop) is forfeited.
49#[derive(Debug)]
50pub(super) struct RefillWaiter {
51    /// How many tokens the refiller has granted this request with. Read by the acquirer
52    /// on each poll to distinguish a grant from a spurious wakeup.
53    ///
54    /// A zero means "not granted yet" where a non-zero value is the grant itself. A
55    /// request is only queued for a non-zero amount.
56    ///
57    /// Taken by the acquirer which then resets it to zero when it emits a permit.
58    granted: AtomicU64,
59    /// The blocked acquirer's task waker. Re-registered on every poll so it is always
60    /// current, and woken by the refiller on grant.
61    waker: AtomicWaker,
62}
63
64impl RefillWaiter {
65    /// Constructor.
66    pub(super) fn new() -> Self {
67        Self {
68            granted: AtomicU64::new(0),
69            waker: AtomicWaker::new(),
70        }
71    }
72
73    /// Return the number of tokens granted to this waiter. `None` means it
74    /// wasn't granted yet.
75    ///
76    /// See [`Self::set_granted`] for the memory ordering.
77    fn granted(&self) -> Option<NonZero<u64>> {
78        NonZero::new(self.granted.load(Ordering::Relaxed))
79    }
80
81    /// Return true iff this waiter was granted permission to use the requested
82    /// bandwidth.
83    fn is_granted(&self) -> bool {
84        self.granted().is_some()
85    }
86
87    /// Prepare this waiter for a new request with the given `waker`.
88    ///
89    /// This must be called before the waiter is enqueued with the refiller.
90    pub(super) fn prepare(&self, waker: &Waker) {
91        // Reset the waiter with this new waker.
92        self.set_granted(0);
93        self.waker.register(waker);
94    }
95
96    /// Grant a number of tokens for this waiter.
97    ///
98    /// The `granted` value is given because it might be clamped so we record what was
99    /// actually granted rather than what was asked.
100    ///
101    /// This function does the atomic work in the proper order the caller doesn't need to
102    /// bother about. The concurrency handling logic is contained.
103    ///
104    /// A `granted` value of 0 is not possible as such value indicate that the permit has
105    /// not been granted yet.
106    pub(super) fn grant(&self, granted: NonZero<u64>) {
107        // A single store publishes both the amount and the fact that we were funded.
108        self.set_granted(granted.get());
109        // We are granted, wake up the waiter!
110        self.wake();
111    }
112
113    /// Register the given `waker` and return true iff this waiter was granted.
114    ///
115    /// Note on the function name, it doesn't return a Poll state but this was done so
116    /// the function semantic carries the intent of the context.
117    pub(super) fn poll_granted(&self, waker: &Waker) -> bool {
118        // Optimization. Before re-registering the waker (see why below), check if we were
119        // granted. This is very cheap to call and avoids the waker registration complexity.
120        // However, it is NOT critical to the lockless state of this pool.
121        if self.is_granted() {
122            return true;
123        }
124
125        // The async contract here is that even if we registered a previous Waker before,
126        // the one given right now is the one that is expected to be woken. Hence, the
127        // register() again.
128        //
129        // We then check again if the tokens were granted because of the documented
130        // `AtomicWaker` pattern that goes like this:
131        //
132        //      sink:     check if granted -> false
133        //      refiller: grant the tokens. waker.wake()
134        //      sink:     register(new waker). return Pending.
135        //
136        // That new waker is never woken up because the grant was done just before hence
137        // why we re-check the granted tokens.
138        self.waker.register(waker);
139        self.is_granted()
140    }
141
142    /// Take the granted tokens out of this waiter leaving it to zero tokens.
143    ///
144    /// This is used by the acquirer to build the permit once granted.
145    pub(super) fn take_granted(&self) -> u64 {
146        // Relaxed is enough. The grant is the value of this
147        // counter, it gates no other memory in the refiller.
148        self.granted.swap(0, Ordering::Relaxed)
149    }
150
151    /// Set atomically the given `val` as the granted value.
152    ///
153    /// # Ordering
154    ///
155    /// The grant must survive the "lost wakeup" race where the refiller grants and wakes
156    /// before the acquirer has (re-)registered its waker:
157    ///
158    /// ```text
159    ///     refiller:  set_granted(n)         // grant
160    ///     refiller:  waker.wake()           // no waker registered yet => wakes nobody
161    ///     acquirer:  waker.register(cx)     // register, too late for the wake above
162    ///     acquirer:  is_granted() -> ???    // must observe the grant or stuck forever
163    /// ```
164    ///
165    /// The acquirer's re-check after `waker.register` must be forced to observe the grant.
166    /// That happens-before is actually provided by the [`AtomicWaker`].
167    ///
168    /// [`Ordering::Relaxed`] suffices because the counter gates no other memory.
169    fn set_granted(&self, val: u64) {
170        self.granted.store(val, Ordering::Relaxed);
171    }
172
173    /// Wake the waker.
174    fn wake(&self) {
175        self.waker.wake();
176    }
177}
178
179/// A bandwidth refiller is in charge of refilling the associated
180/// [`super::BandwidthPool`] and processing any pending RefillWaiter that were enqueued
181/// by the pool.
182///
183/// There is exactly one refiller per pool as it owns the receiving end of the request
184/// channel. The channel is a FIFO of waiters.
185///
186/// Dropping it closes the pool which makes queued and new acquirers fail with
187/// [`super::BwPoolError::PoolClosed`].
188#[derive(Debug)]
189pub struct BandwidthRefiller {
190    /// The shared token bucket which comes from the [`super::BandwidthPool`].
191    bucket: Arc<AtomicTokenBucket>,
192    /// Receiving end of the request channel
193    rx: mpsc::UnboundedReceiver<BwRequest>,
194    /// A single request we have taken off the channel to inspect but cannot yet fund. If
195    /// only mpsc channels had an "is_empty()".
196    ///
197    /// This is populated by the [`Self::wait`] function
198    head: Option<BwRequest>,
199    /// Tokens that have been drained out of the fast-path bucket or handed in via
200    /// [`Self::refill_and_serve`] but not yet distributed.
201    ///
202    /// If anything is left at the end of the refill loop, it is published back in the
203    /// main pool fast path.
204    held: u64,
205}
206
207impl BandwidthRefiller {
208    /// Constructor.
209    pub(super) fn new(
210        bucket: Arc<AtomicTokenBucket>,
211        rx: mpsc::UnboundedReceiver<BwRequest>,
212    ) -> Self {
213        Self {
214            bucket,
215            rx,
216            head: None,
217            held: 0,
218        }
219    }
220
221    /// Start the refiller main loop. This should be run in its own task as it is
222    /// blocking until the bandwidth pool closes.
223    ///
224    /// Using the given [`SleepProvider`], we drive the refill with it along side a
225    /// [`TokenBucket`] that is built at the start with the given `config`.
226    ///
227    /// This waits on new request that comes in when the fast-path is depleted that is a
228    /// request waiting for a refill. Once a request is received, a refill is triggered
229    /// and then the loop sleeps until the needed deficit is available.
230    ///
231    /// As an example, if the queue has a request for 10 tokens but only 5 are available
232    /// in the pool after an immediate refill, we will sleep the exact time it takes to
233    /// get another 5 tokens. Then, it goes on to the next request and so on until the
234    /// queue is empty.
235    ///
236    /// Keen observer will notice that once a request is received, the refiller will have
237    /// to empty the entire queue before the fast path could even see 1 token added back
238    /// by a refill. That is because each iteration of the loop sleeps the exact amount
239    /// of time to fulfill the pending request.
240    ///
241    /// Returns when the pool is closed or if the config rate is zero or if the requested
242    /// amount of token is above the pool capacity.
243    pub async fn run<SP: SleepProvider>(mut self, sleep: SP, rate: u64) {
244        let config = TokenBucketConfig {
245            rate,
246            bucket_max: self.bucket.capacity(),
247        };
248        let mut bucket = TokenBucket::new(&config, sleep.now());
249        // Start empty. The pool's first start full. This avoids adding a second burst to
250        // the pool after the fast path is depleted.
251        let _ = bucket.drain_all();
252
253        loop {
254            // Wait on the "doorbell" that is a pending request wanting tokens.
255            if !self.wait().await {
256                // The pool is closed.
257                return;
258            }
259
260            // Serve the queue. Sleep for the deficit until queue is empty.
261            loop {
262                // Refill the bucket with what we can.
263                bucket.refill(sleep.now());
264                // This will reserve all available tokens from the pool so the fast-path
265                // ends up empty and then starts serving queued requests. We use all
266                // token from the TokenBucket as well.
267                //
268                // If everyone is served and we still have tokens, they are put back
269                // in the fast path.
270                match self.refill_and_serve(bucket.drain_all()) {
271                    // We served everyone, fast-path has the surplus if any. Go back to
272                    // the doorbell.
273                    None => break,
274                    // The pending request wants `deficit` amount of tokens, wait for that
275                    // exact value.
276                    Some(deficit) => match bucket.tokens_available_at(deficit) {
277                        Ok(at) => {
278                            let d = at.saturating_duration_since(sleep.now());
279                            sleep.sleep(d).await;
280                        }
281                        // Zero rate or a deficit above the burst. We can never fulfill
282                        // that request so error else we are stuck.
283                        Err(_) => return,
284                    },
285                }
286            }
287        }
288    }
289
290    /// Wait until at least one bandwidth request is queued.
291    ///
292    /// Returns `true` once a request is received or `false` if the pool has been closed
293    /// meaning the tx end is closed.
294    ///
295    /// Returns `true` immediately if a request is already held as the head.
296    ///
297    /// This should only be used as a "doorbell" that is indicating someone is at the
298    /// door with a request rather than waiting for the next request. One should use
299    /// `Self::serve` for that.
300    #[cfg_attr(feature = "bench", visibility::make(pub))]
301    pub(crate) async fn wait(&mut self) -> bool {
302        // Avoid overwriting an existing request.
303        if self.head.is_some() {
304            return true;
305        }
306        match self.rx.next().await {
307            Some(req) => {
308                self.head = Some(req);
309                true
310            }
311            // TODO(relay): Need to fix this with partial permit else this will make the
312            // whole bucket fail if a request above capacity shows up.
313            None => false,
314        }
315    }
316
317    /// Add `tokens` to the pool and then serve all pending requests if any.
318    ///
319    /// The very first thing that this function does is drain the pool's fast path tokens
320    /// in order to reserve its current balance for the waiting queue.
321    ///
322    /// Any surplus left will be put back into the pool's fast path.
323    ///
324    /// Returns `None` if no one is left waiting indicating the pool is now idle and
325    /// [`Self::wait`] can be safely used to get notified of a new request.
326    ///
327    /// Returns `Some(deficit)` if an acquirer is still waiting where `deficit` is how
328    /// many more tokens are needed before it can be served. The caller can use this to
329    /// decide how long to wait before the next refill.
330    #[cfg_attr(feature = "bench", visibility::make(pub))]
331    pub(crate) fn refill_and_serve(&mut self, tokens: u64) -> Option<u64> {
332        let capacity = self.bucket.capacity();
333
334        // Reclaim tokens sitting in the fast path.
335        let reclaimed = self.bucket.drain();
336        self.held = self
337            .held
338            .saturating_add(reclaimed)
339            .saturating_add(tokens)
340            .min(capacity);
341
342        // Serve the request queue with the token we are holding.
343        self.serve(capacity);
344
345        // If we still have a head, report its deficit, else publish the remaining tokens
346        // in the pool's fast path. We use the snapshot capacity here so it is the same
347        // value used for the serve.
348        //
349        // The clamp matters as an unclamped deficit could be above the burst and our
350        // caller would then never be able to wait for it.
351        match &self.head {
352            Some((needed, _)) => Some((*needed).min(capacity).saturating_sub(self.held)),
353            None => {
354                self.publish_held();
355                None
356            }
357        }
358    }
359
360    /// Serve pending requests with the token we are holding.
361    ///
362    /// The given `capacity` is essentially the maximum we can give a single request.
363    ///
364    /// If we have a token deficit, the head is updated with the latest request that we
365    /// can't serve which indicates the caller we are in deficit.
366    fn serve(&mut self, capacity: u64) {
367        loop {
368            let (wanted, waiter) = match self.head.take() {
369                Some(req) => req,
370                None => match self.rx.try_recv() {
371                    Ok(req) => req,
372                    // Channel is empty or closed. We are done.
373                    Err(_) => return,
374                },
375            };
376
377            // DO NOT REMOVE THIS
378            //
379            // A queued request is sent as asked so this clamp is the only thing capping
380            // the slow path. We can never hold more than the burst and so a larger
381            // request would sit in the queue forever.
382            let needed = wanted.min(capacity);
383            if needed > self.held {
384                // Unable to permit this request, keep it for next round.
385                self.head = Some((wanted, waiter));
386                return;
387            }
388
389            // Commit the permit first and then wake. If the waker (acquirer) was torn
390            // down, the grant is forfeited but that is a documented limitation.
391            self.held -= needed;
392            // Grant what we took which could be a clamped value.
393            waiter.grant(NonZero::new(needed).expect("BW request is 0 tokens"));
394        }
395    }
396
397    /// Publish any held surplus back to the fast-path.
398    fn publish_held(&mut self) {
399        self.bucket.refill(self.held);
400        self.held = 0;
401    }
402}
403
404impl Drop for BandwidthRefiller {
405    /// Wake all queued waiters on teardown so they wake up and get to realize the pool
406    /// is closed on their next poll.
407    fn drop(&mut self) {
408        // Close the receiver so any new waiter gets a pool closed error.
409        self.rx.close();
410        // The waiter we pulled off the channel as the head but never served.
411        if let Some((_, head)) = self.head.take() {
412            head.wake();
413        }
414        // Wake any enqueued waiters.
415        while let Ok((_, waiter)) = self.rx.try_recv() {
416            waiter.wake();
417        }
418    }
419}
420
421#[cfg(test)]
422mod test {
423    // @@ begin test lint list maintained by maint/add_warning @@
424    #![allow(clippy::bool_assert_comparison)]
425    #![allow(clippy::clone_on_copy)]
426    #![allow(clippy::dbg_macro)]
427    #![allow(clippy::mixed_attributes_style)]
428    #![allow(clippy::print_stderr)]
429    #![allow(clippy::print_stdout)]
430    #![allow(clippy::single_char_pattern)]
431    #![allow(clippy::unwrap_used)]
432    #![allow(clippy::unchecked_time_subtraction)]
433    #![allow(clippy::useless_vec)]
434    #![allow(clippy::needless_pass_by_value)]
435    #![allow(clippy::string_slice)] // See arti#2571
436    //! <!-- @@ end test lint list maintained by maint/add_warning @@ -->
437
438    use super::*;
439
440    use futures::FutureExt as _;
441    use futures::task::{ArcWake, waker};
442    use std::sync::atomic::AtomicUsize;
443
444    /// Build a new drained refiller of `capacity` and the request channel sender used to
445    /// enqueue requests.
446    fn drained_refiller(capacity: u64) -> (mpsc::UnboundedSender<BwRequest>, BandwidthRefiller) {
447        let (tx, rx) = mpsc::unbounded();
448        let bucket = Arc::new(AtomicTokenBucket::new(capacity));
449        assert_eq!(bucket.claim(capacity), Some(capacity)); // the bucket starts full; empty it
450        (tx, BandwidthRefiller::new(bucket, rx))
451    }
452
453    /// Enqueue a request for `needed` tokens.
454    ///
455    /// Return its waiter so the test can observe the grant.
456    fn enqueue(tx: &mpsc::UnboundedSender<BwRequest>, needed: u64) -> Arc<RefillWaiter> {
457        let waiter = Arc::new(RefillWaiter::new());
458        tx.unbounded_send((needed, Arc::clone(&waiter))).unwrap();
459        waiter
460    }
461
462    /// A waker that counts how many times it is woken.
463    #[derive(Default)]
464    struct WakeCount(AtomicUsize);
465
466    impl ArcWake for WakeCount {
467        fn wake_by_ref(arc_self: &Arc<Self>) {
468            arc_self.0.fetch_add(1, Ordering::SeqCst);
469        }
470    }
471
472    #[test]
473    fn deficit() {
474        let (tx, mut r) = drained_refiller(100);
475        let w = enqueue(&tx, 50);
476
477        // Partial refills which report the shrinking deficit. No grant as we don't have
478        // enough.
479        assert_eq!(r.refill_and_serve(20), Some(30));
480        assert_eq!(r.refill_and_serve(20), Some(10));
481        assert!(!w.is_granted());
482
483        // Last refill before reaching what is needed.
484        assert_eq!(r.refill_and_serve(10), None);
485        assert!(w.is_granted());
486    }
487
488    #[test]
489    fn serve_fifo() {
490        let (tx, mut r) = drained_refiller(100);
491        // Enqueue two acquirers.
492        let a = enqueue(&tx, 30);
493        let b = enqueue(&tx, 30);
494
495        // Refill with 40 tokens, A should be granted wanting 30.
496        assert_eq!(r.refill_and_serve(40), Some(20));
497        assert!(a.is_granted());
498        // Not granted, 10 remains for B with a 20 deficit.
499        assert!(!b.is_granted());
500        assert_eq!(r.bucket.available(), 0);
501        // Refill the deficit and B should be granted.
502        assert_eq!(r.refill_and_serve(20), None);
503        assert!(b.is_granted());
504    }
505
506    #[test]
507    fn reclaim_fast_path() {
508        // The bucket holds 40 tokens that a failed fast-path attempt for 50 could not claim.
509        let (tx, rx) = mpsc::unbounded();
510        let bucket = Arc::new(AtomicTokenBucket::new(100));
511        assert_eq!(bucket.claim(60), Some(60));
512        // Pool has 40 now. Enqueue a request for 50.
513        let mut r = BandwidthRefiller::new(Arc::clone(&bucket), rx);
514        let w = enqueue(&tx, 50);
515
516        // A refill of 0 should take those 40 from the fast path and put them in the
517        // refiller held reserve returning a deficit of 10 to grant the request of 50.
518        assert_eq!(r.refill_and_serve(0), Some(10));
519        assert_eq!(bucket.available(), 0);
520        // Refill 20 more, the request should be granted and 10 should be put in the fast
521        // path.
522        assert_eq!(r.refill_and_serve(20), None);
523        assert!(w.is_granted());
524        assert_eq!(bucket.available(), 10);
525    }
526
527    #[test]
528    fn wake_on_permit() {
529        let (tx, mut r) = drained_refiller(100);
530        let w = enqueue(&tx, 50);
531        let wake_counter = Arc::new(WakeCount::default());
532        w.waker.register(&waker(Arc::clone(&wake_counter)));
533
534        // Not enough to wake the waiter. Request wants 50 so deficit is now 30.
535        assert_eq!(r.refill_and_serve(20), Some(30));
536        assert_eq!(wake_counter.0.load(Ordering::SeqCst), 0);
537        // Refills with the deficit, the waker should wake up and be granted.
538        assert_eq!(r.refill_and_serve(30), None);
539        assert_eq!(wake_counter.0.load(Ordering::SeqCst), 1);
540        assert!(w.is_granted());
541    }
542
543    #[test]
544    fn wait_doorbell() {
545        let (tx, mut r) = drained_refiller(100);
546
547        // Nothing queued: wait() pends.
548        assert_eq!(r.wait().now_or_never(), None);
549
550        // A request rings the doorbell; it is taken as the head and served by
551        // the following refill.
552        let w = enqueue(&tx, 10);
553        assert_eq!(r.wait().now_or_never(), Some(true));
554        assert_eq!(r.refill_and_serve(10), None);
555        assert!(w.is_granted());
556        // All senders gone, wait() should report that nothing is there.
557        drop(tx);
558        assert_eq!(r.wait().now_or_never(), Some(false));
559    }
560}