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}