Skip to main content

tor_async_utils/global_rate_limit/
reader.rs

1//! A rate limited async reader.
2//!
3//! An [`AsyncRead`] wrapper that rate limits the bytes it yields by acquiring bandwidth
4//! from a [`crate::bw_pool::BandwidthPool`] before each read. One byte is one token.
5//!
6//! # Example
7//!
8//! ```
9//! use std::num::NonZero;
10//! use futures::AsyncReadExt as _;
11//! use tor_async_utils::global_rate_limit::GlobalRateLimitedReader;
12//! use tor_async_utils::bw_pool::BandwidthPool;
13//!
14//! futures::executor::block_on(async {
15//!     let (pool, _refiller) = BandwidthPool::new(64 * 1024);
16//!     let max_chunk = NonZero::new(1024).unwrap();
17//!     let mut reader =
18//!         GlobalRateLimitedReader::new(&b"hello world"[..], pool.new_acquirer(), max_chunk);
19//!
20//!     // The pool starts full so this is served from the fast path.
21//!     let mut buf = [0; 10];
22//!     reader.read_exact(&mut buf).await.unwrap();
23//! });
24//! ```
25
26use futures::AsyncRead;
27use pin_project::pin_project;
28use std::io::Error;
29use std::num::NonZero;
30use std::pin::Pin;
31use std::task::{Context, Poll, ready};
32
33use super::DirectionState;
34use crate::bw_pool::BandwidthAcquirer;
35
36/// An [`AsyncRead`] wrapper that acquires bandwidth before yielding bytes.
37///
38/// Every byte read costs one token from the shared pool. A single
39/// [`AsyncRead::poll_read`] is capped to the pool's bandwidth so a large buffer is
40/// filled in many chunks.
41#[derive(Debug)]
42#[pin_project]
43pub struct GlobalRateLimitedReader<R> {
44    /// The underlying reader bytes are read from.
45    #[pin]
46    inner: R,
47    /// The per-direction state holding an acquirer and permit.
48    state: DirectionState,
49}
50
51impl<R> GlobalRateLimitedReader<R> {
52    /// Construct a reader that spends one token per byte from the `acquirer`'s pool.
53    ///
54    /// Each read requests at most `max_chunk` tokens.
55    pub fn new(inner: R, acquirer: BandwidthAcquirer, max_chunk: NonZero<usize>) -> Self {
56        Self {
57            inner,
58            state: DirectionState::new(acquirer, max_chunk),
59        }
60    }
61}
62
63impl<R> AsyncRead for GlobalRateLimitedReader<R>
64where
65    R: AsyncRead,
66{
67    fn poll_read(
68        self: Pin<&mut Self>,
69        cx: &mut Context<'_>,
70        buf: &mut [u8],
71    ) -> Poll<Result<usize, Error>> {
72        let this = self.project();
73        poll_read_limited(this.inner, this.state, cx, buf)
74    }
75}
76
77/// Helper: Rate-limited [`AsyncRead::poll_read`] of `inner` using the given read
78/// direction `state`. This is used by multiple object hence why in a read helper.
79pub(super) fn poll_read_limited<R: AsyncRead>(
80    inner: Pin<&mut R>,
81    state: &mut DirectionState,
82    cx: &mut Context<'_>,
83    buf: &mut [u8],
84) -> Poll<Result<usize, Error>> {
85    // For an empty buffer, just defer to the inner, no need to bother for a permit.
86    if buf.is_empty() {
87        return inner.poll_read(cx, buf);
88    }
89
90    // Acquire a permit and learn how many bytes we are cleared to read.
91    let permit = ready!(state.poll_acquire(cx, buf.len()))?;
92    let len = super::to_usize(permit.granted().min(super::to_u64(buf.len())));
93    let buf = &mut buf[..len];
94
95    match inner.poll_read(cx, buf) {
96        // The inner is not ready. Drop the permit so the tokens go back to the pool for
97        // someone else to use rather than being parked here until this read works. We
98        // acquire again on the next poll.
99        Poll::Pending => Poll::Pending,
100        // The inner had an error, drop the permit to refund.
101        Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
102        // Claim what was read and refund the rest.
103        Poll::Ready(Ok(read)) => {
104            DirectionState::commit(permit, read);
105            Poll::Ready(Ok(read))
106        }
107    }
108}
109
110#[cfg(test)]
111mod test {
112    // @@ begin test lint list maintained by maint/add_warning @@
113    #![allow(clippy::bool_assert_comparison)]
114    #![allow(clippy::clone_on_copy)]
115    #![allow(clippy::dbg_macro)]
116    #![allow(clippy::mixed_attributes_style)]
117    #![allow(clippy::print_stderr)]
118    #![allow(clippy::print_stdout)]
119    #![allow(clippy::single_char_pattern)]
120    #![allow(clippy::unwrap_used)]
121    #![allow(clippy::unchecked_time_subtraction)]
122    #![allow(clippy::useless_vec)]
123    #![allow(clippy::needless_pass_by_value)]
124    #![allow(clippy::string_slice)] // See arti#2571
125    //! <!-- @@ end test lint list maintained by maint/add_warning @@ -->
126
127    use super::*;
128
129    use futures::{AsyncReadExt as _, FutureExt as _};
130
131    use crate::bw_pool::BandwidthPool;
132
133    /// Max chunk used by the tests that don't exercise the cap itself.
134    const MAX_CHUNK: NonZero<usize> = NonZero::new(1024).unwrap();
135
136    #[test]
137    fn fast_path() {
138        let (pool, _refiller) = BandwidthPool::new(100);
139        let mut reader =
140            GlobalRateLimitedReader::new(&[1_u8; 50][..], pool.new_acquirer(), MAX_CHUNK);
141
142        // Read 30 bytes hits the fast path.
143        let mut buf = [0; 30];
144        assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 30);
145        assert_eq!(pool.available(), 70);
146        assert_eq!(buf, [1; 30]);
147    }
148
149    #[test]
150    fn capped_pool_capacity() {
151        let (pool, _refiller) = BandwidthPool::new(30);
152        let mut reader =
153            GlobalRateLimitedReader::new(&[1_u8; 100][..], pool.new_acquirer(), MAX_CHUNK);
154
155        // Read 100 in a pool of capacity 30 means only 30 is written.
156        let mut buf = [0; 100];
157        assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 30);
158        assert_eq!(pool.available(), 0);
159    }
160
161    #[test]
162    fn max_chunk() {
163        let (pool, _refiller) = BandwidthPool::new(100);
164        let cap = NonZero::new(10).unwrap();
165        let mut reader = GlobalRateLimitedReader::new(&[1_u8; 50][..], pool.new_acquirer(), cap);
166
167        // Buffer is 30 but max_chunk caps the request to 10 tokens.
168        let mut buf = [0; 30];
169        assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 10);
170        assert_eq!(pool.available(), 90);
171    }
172
173    #[test]
174    fn pending() {
175        let (pool, mut refiller) = BandwidthPool::new(30);
176        let mut reader =
177            GlobalRateLimitedReader::new(&[1_u8; 100][..], pool.new_acquirer(), MAX_CHUNK);
178
179        // Empty the pool with a read of 30.
180        let mut buf = [0; 30];
181        assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 30);
182
183        // Pool is empty so the next read is Pending until a refill.
184        let mut read = reader.read(&mut buf);
185        assert!((&mut read).now_or_never().is_none());
186        assert_eq!(refiller.refill_and_serve(30), None);
187        assert_eq!((&mut read).now_or_never().unwrap().unwrap(), 30);
188    }
189
190    #[test]
191    fn pool_closed() {
192        let (pool, refiller) = BandwidthPool::new(30);
193        let mut reader =
194            GlobalRateLimitedReader::new(&[1_u8; 200][..], pool.new_acquirer(), MAX_CHUNK);
195
196        // Read 10.
197        let mut buf = [0; 10];
198        assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 10);
199
200        // Drain the pool then drop the refiller. The next read has to enqueue but it
201        // will fail because the pool has been closed due to the refiller closing.
202        let mut buf = [0; 20];
203        assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 20);
204        drop(refiller);
205        let mut buf = [0; 10];
206        assert!(reader.read(&mut buf).now_or_never().unwrap().is_err());
207    }
208}