Skip to main content

tor_async_utils/global_rate_limit/
writer.rs

1//! A rate limited async writer.
2//!
3//! An [`AsyncWrite`] wrapper that rate limits the bytes it forwards by acquiring bandwidth
4//! from a [`crate::bw_pool::BandwidthPool`] before each write. One byte is one token.
5//!
6//! # Example
7//!
8//! ```
9//! use std::num::NonZero;
10//! use futures::AsyncWriteExt as _;
11//! use tor_async_utils::global_rate_limit::GlobalRateLimitedWriter;
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 writer = GlobalRateLimitedWriter::new(Vec::new(), pool.new_acquirer(), max_chunk);
18//!
19//!     // The pool starts full so this is served from the fast path.
20//!     writer.write_all(&[0; 512]).await.unwrap();
21//! });
22//! ```
23
24use futures::AsyncWrite;
25use pin_project::pin_project;
26use std::io::Error;
27use std::num::NonZero;
28use std::pin::Pin;
29use std::task::{Context, Poll, ready};
30
31use super::DirectionState;
32use crate::bw_pool::BandwidthAcquirer;
33
34/// An [`AsyncWrite`] wrapper that acquires bandwidth before writing bytes.
35///
36/// A single byte is one token which we acquire from the shared pool. A single
37/// [`AsyncWrite::poll_write`] is capped to the pool's capacity so a large buffer has to
38/// go in written several chunks.
39#[derive(Debug)]
40#[pin_project]
41pub struct GlobalRateLimitedWriter<W> {
42    /// The underlying writer bytes are forwarded to.
43    #[pin]
44    inner: W,
45    /// The per-direction state holding an acquirer and permit.
46    state: DirectionState,
47}
48
49impl<W> GlobalRateLimitedWriter<W> {
50    /// Constructor.
51    ///
52    /// This writer is rate limited as 1 byte per token.
53    ///
54    /// Each write requests at most `max_chunk` tokens.
55    pub fn new(inner: W, acquirer: BandwidthAcquirer, max_chunk: NonZero<usize>) -> Self {
56        Self {
57            inner,
58            state: DirectionState::new(acquirer, max_chunk),
59        }
60    }
61}
62
63impl<W> AsyncWrite for GlobalRateLimitedWriter<W>
64where
65    W: AsyncWrite,
66{
67    fn poll_write(
68        self: Pin<&mut Self>,
69        cx: &mut Context<'_>,
70        buf: &[u8],
71    ) -> Poll<Result<usize, Error>> {
72        let this = self.project();
73        poll_write_limited(this.inner, this.state, cx, buf)
74    }
75
76    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
77        self.project().inner.poll_flush(cx)
78    }
79
80    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
81        self.project().inner.poll_close(cx)
82    }
83}
84
85/// Helper: Rate-limited [`AsyncWrite::poll_write`] of `inner` using the given read
86/// direction `state`. This is used by multiple object hence why in a write helper.
87pub(super) fn poll_write_limited<W: AsyncWrite>(
88    inner: Pin<&mut W>,
89    state: &mut DirectionState,
90    cx: &mut Context<'_>,
91    buf: &[u8],
92) -> Poll<Result<usize, Error>> {
93    // For an empty buffer, just defer to the inner, no need to bother for a permit.
94    if buf.is_empty() {
95        return inner.poll_write(cx, buf);
96    }
97
98    // Acquire a permit and get the permit granted tokens worth of bytes from the buffer.
99    let permit = ready!(state.poll_acquire(cx, buf.len()))?;
100    let len = super::to_usize(permit.granted().min(super::to_u64(buf.len())));
101    let buf = &buf[..len];
102
103    match inner.poll_write(cx, buf) {
104        // The inner is not ready, refund the permit by dropping it.
105        Poll::Pending => Poll::Pending,
106        // The inner had an error, refund the permit by dropping it.
107        Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
108        // Claim what was sent and refund the rest.
109        Poll::Ready(Ok(written)) => {
110            DirectionState::commit(permit, written);
111            Poll::Ready(Ok(written))
112        }
113    }
114}
115
116#[cfg(test)]
117mod test {
118    // @@ begin test lint list maintained by maint/add_warning @@
119    #![allow(clippy::bool_assert_comparison)]
120    #![allow(clippy::clone_on_copy)]
121    #![allow(clippy::dbg_macro)]
122    #![allow(clippy::mixed_attributes_style)]
123    #![allow(clippy::print_stderr)]
124    #![allow(clippy::print_stdout)]
125    #![allow(clippy::single_char_pattern)]
126    #![allow(clippy::unwrap_used)]
127    #![allow(clippy::unchecked_time_subtraction)]
128    #![allow(clippy::useless_vec)]
129    #![allow(clippy::needless_pass_by_value)]
130    #![allow(clippy::string_slice)] // See arti#2571
131    //! <!-- @@ end test lint list maintained by maint/add_warning @@ -->
132
133    use super::*;
134
135    use futures::{AsyncWriteExt as _, FutureExt as _};
136
137    use crate::bw_pool::BandwidthPool;
138
139    /// Max chunk used by the tests that don't exercise the cap itself.
140    const MAX_CHUNK: NonZero<usize> = NonZero::new(1024).unwrap();
141
142    /// An [`AsyncWrite`] that is never ready to take bytes. Needed to test a stalled
143    /// connection as to test the refund of dropping permit when pending.
144    struct NeverReady;
145
146    impl AsyncWrite for NeverReady {
147        fn poll_write(
148            self: Pin<&mut Self>,
149            _cx: &mut Context<'_>,
150            _buf: &[u8],
151        ) -> Poll<Result<usize, Error>> {
152            Poll::Pending
153        }
154        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
155            Poll::Ready(Ok(()))
156        }
157        fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
158            Poll::Ready(Ok(()))
159        }
160    }
161
162    #[test]
163    fn inner_pending_refunds() {
164        let (pool, _refiller) = BandwidthPool::new(100);
165        let mut writer = GlobalRateLimitedWriter::new(NeverReady, pool.new_acquirer(), MAX_CHUNK);
166
167        // The permit is acquired but the inner can't take the bytes.
168        assert!(writer.write(&[0; 30]).now_or_never().is_none());
169
170        // The permit was dropped rather than parked so every token is back in the pool
171        // for a connection that can actually use it.
172        assert_eq!(pool.available(), 100);
173
174        // Same on the next poll, nothing accumulates.
175        assert!(writer.write(&[0; 30]).now_or_never().is_none());
176        assert_eq!(pool.available(), 100);
177    }
178
179    #[test]
180    fn fast_path() {
181        let (pool, _refiller) = BandwidthPool::new(100);
182        let mut writer = GlobalRateLimitedWriter::new(Vec::new(), pool.new_acquirer(), MAX_CHUNK);
183
184        // Writing 30 bytes hits the fast path.
185        assert_eq!(writer.write(&[0; 30]).now_or_never().unwrap().unwrap(), 30);
186        assert_eq!(pool.available(), 70);
187    }
188
189    #[test]
190    fn capped_pool_capacity() {
191        let (pool, _refiller) = BandwidthPool::new(30);
192        let mut writer = GlobalRateLimitedWriter::new(Vec::new(), pool.new_acquirer(), MAX_CHUNK);
193
194        // Writing 100 in a pool of capacity 30 means only 30 is written.
195        assert_eq!(writer.write(&[0; 100]).now_or_never().unwrap().unwrap(), 30);
196        assert_eq!(pool.available(), 0);
197    }
198
199    #[test]
200    fn max_chunk() {
201        let (pool, _refiller) = BandwidthPool::new(100);
202        let cap = NonZero::new(10).unwrap();
203        let mut writer = GlobalRateLimitedWriter::new(Vec::new(), pool.new_acquirer(), cap);
204
205        // Buffer is 30 but max_chunk caps the request to 10 tokens.
206        assert_eq!(writer.write(&[0; 30]).now_or_never().unwrap().unwrap(), 10);
207        assert_eq!(pool.available(), 90);
208    }
209
210    #[test]
211    fn pending() {
212        let (pool, mut refiller) = BandwidthPool::new(30);
213        let mut writer = GlobalRateLimitedWriter::new(Vec::new(), pool.new_acquirer(), MAX_CHUNK);
214
215        // Empty the pool with a write of 30.
216        assert_eq!(writer.write(&[0; 30]).now_or_never().unwrap().unwrap(), 30);
217
218        // Pool is empty so the next write is Pending until a refill.
219        let mut write = writer.write(&[0; 30]);
220        assert!((&mut write).now_or_never().is_none());
221        assert_eq!(refiller.refill_and_serve(30), None);
222        // Pool is refilled, 30 is written.
223        assert_eq!((&mut write).now_or_never().unwrap().unwrap(), 30);
224    }
225
226    #[test]
227    fn pool_closed() {
228        let (pool, refiller) = BandwidthPool::new(30);
229        let mut writer = GlobalRateLimitedWriter::new(Vec::new(), pool.new_acquirer(), MAX_CHUNK);
230
231        // Write 10.
232        assert_eq!(writer.write(&[0; 10]).now_or_never().unwrap().unwrap(), 10);
233
234        // Drain the pool then drop the refiller. The next write has to enqueue but it
235        // will fail because the pool has been closed due to the refiller closing.
236        assert_eq!(writer.write(&[0; 20]).now_or_never().unwrap().unwrap(), 20);
237        drop(refiller);
238        assert!(writer.write(&[0; 10]).now_or_never().unwrap().is_err());
239    }
240}