Skip to main content

tor_async_utils/
global_rate_limit.rs

1//! Rate limiter objects backed by a shared [`crate::bw_pool::BandwidthPool`] making them
2//! global as in sharable accross multiple thread/tasks.
3//!
4//! Available in this module is:
5//!     * A [`sink::GlobalRateLimitedSink`] that implements [`futures::Sink`].
6//!     * A [`writer::GlobalRateLimitedWriter`] that implements [`futures::AsyncWrite`].
7//!     * A [`reader::GlobalRateLimitedReader`] that implements [`futures::AsyncRead`].
8//!     * A [`conn::GlobalRateLimitedConn`] that implements both [`futures::AsyncRead`] and
9//!       [`futures::AsyncWrite`], rate limiting each direction independently.
10//!
11//! Please read carefully each submodule documentation before using. These can be tricky
12//! to operate without a license ;).
13
14mod conn;
15mod reader;
16mod sink;
17mod writer;
18
19pub use conn::GlobalRateLimitedConn;
20pub use reader::GlobalRateLimitedReader;
21pub use sink::GlobalRateLimitedSink;
22pub use writer::GlobalRateLimitedWriter;
23
24use std::io::Error;
25use std::num::NonZero;
26use std::task::{Context, Poll, ready};
27
28use crate::bw_pool::{BandwidthAcquirer, Permit};
29
30/// Convert a `usize` to `u64`. Infallible on every platform we support.
31fn to_u64(x: usize) -> u64 {
32    x.try_into().expect("failed usize to u64 conversion")
33}
34
35/// Convert a `u64` to `usize`. Panics if the value exceeds `usize::MAX`.
36fn to_usize(x: u64) -> usize {
37    x.try_into().expect("failed u64 to usize conversion")
38}
39
40/// Rate-limiting state for a single direction.
41///
42/// It has everything needed to rate limit one direction that is a [`BandwidthAcquirer`]
43/// and an optional maximum chunk.
44///
45/// This is used by the [`GlobalRateLimitedReader`] and [`GlobalRateLimitedWriter`] as
46/// they share that same behavior for each direction (read and write).
47#[derive(Debug)]
48struct DirectionState {
49    /// Acquirer used to get a [`Permit`] from the pool for each poll.
50    acquirer: BandwidthAcquirer,
51    /// Cap on how many tokens a single IO can request.
52    ///
53    /// An IO requests at most this many tokens as long as the buffer is bigger.
54    max_chunk: NonZero<usize>,
55}
56
57impl DirectionState {
58    /// Constructor.
59    fn new(acquirer: BandwidthAcquirer, max_chunk: NonZero<usize>) -> Self {
60        Self {
61            acquirer,
62            max_chunk,
63        }
64    }
65
66    /// Acquire a permit for an IO of the given amount of `tokens`.
67    ///
68    /// The request is capped to the state's max chunk. The grant itself is capped
69    /// to the pool capacity so the returned [`Permit`] can hold less than `tokens`.
70    ///
71    /// The [`Permit`] is handed to the caller rather than kept here so that it lives
72    /// exactly as long as the IO attempt it is for. If the underlying IO turns out to be
73    /// [`Poll::Pending`], the caller drops it and the tokens are refunded to the pool
74    /// instead of being parked for as long as the connection is not ready. That matters
75    /// with many connections as we don't want to hold off ready connections on already
76    /// allocated permits for non ready connections.
77    ///
78    /// The cost is that the caller re-acquires on the next poll rather than resuming
79    /// with what it already had. It is by design.
80    fn poll_acquire(&mut self, cx: &mut Context<'_>, tokens: usize) -> Poll<Result<Permit, Error>> {
81        let want = tokens.min(self.max_chunk.get());
82        let permit = ready!(self.acquirer.poll_acquire(cx, to_u64(want))).map_err(Error::other)?;
83        Poll::Ready(Ok(permit))
84    }
85
86    /// Claim `tokens` on `permit` after a successful IO.
87    ///
88    /// `tokens` must not exceed what the permit was granted. Callers get that by capping
89    /// the IO to [`Permit::granted`].
90    ///
91    /// The permit is consumed so whatever is left unclaimed is refunded on drop.
92    fn commit(mut permit: Permit, tokens: usize) {
93        let claimed = permit.claim(to_u64(tokens));
94        debug_assert!(claimed.is_ok(), "IO reported more than the permit granted");
95        if claimed.is_err() {
96            // The inner misbehaved. Claim it all to avoid refunding what was used.
97            permit.claim_all();
98        }
99    }
100}
101
102/// Error returned by the rate limiters in this module.
103#[derive(Debug, thiserror::Error)]
104#[non_exhaustive]
105pub enum GlobalRateLimitedError<E> {
106    /// The bandwidth pool is on error. No more refiller.
107    #[error("bandwidth pool error")]
108    Pool(#[from] crate::bw_pool::BwPoolError),
109    /// The underlying sink failed.
110    #[error("underlying sink error")]
111    Sink(#[source] E),
112    /// No permit
113    #[error("no permit when sending")]
114    MissingPermit,
115}