Skip to main content

tor_async_utils/global_rate_limit/
conn.rs

1//! A rate limited async duplex connection.
2//!
3//! An [`AsyncRead`] and [`AsyncWrite`] wrapper that rate limits both directions of a
4//! single bidirectional byte stream such as a TCP or TLS connection, by acquiring
5//! bandwidth from a [`super::BandwidthAcquirer`] given for each diection (read and
6//! write). One byte is one token.
7
8use futures::{AsyncRead, AsyncWrite};
9use pin_project::pin_project;
10use std::io::Error;
11use std::num::NonZero;
12use std::pin::Pin;
13use std::task::{Context, Poll};
14
15use tor_rtcompat::StreamOps;
16
17use super::DirectionState;
18use crate::bw_pool::BandwidthAcquirer;
19
20/// An [`AsyncRead`] and [`AsyncWrite`] wrapper that rate limits each direction of a single
21/// bidirectional byte stream.
22///
23/// Every byte read or written costs one token per direction pool.
24#[derive(Debug)]
25#[pin_project]
26pub struct GlobalRateLimitedConn<S> {
27    /// The underlying stream.
28    #[pin]
29    inner: S,
30    /// Rate-limiting state for the read direction.
31    read_state: DirectionState,
32    /// Rate-limiting state for the write direction.
33    write_state: DirectionState,
34}
35
36impl<S> GlobalRateLimitedConn<S> {
37    /// Constructor.
38    ///
39    /// We recommend that the `read_acquirer` and `write_acquirer` comme from different
40    /// bandwidth pools so one direction doesn't starve the other side. In a
41    /// bidirectional setup, this could be equivalent to unidirectionnal.
42    ///
43    /// Each IO requests at most `max_chunk` tokens in either direction.
44    pub fn new(
45        inner: S,
46        read_acquirer: BandwidthAcquirer,
47        write_acquirer: BandwidthAcquirer,
48        max_chunk: NonZero<usize>,
49    ) -> Self {
50        Self {
51            inner,
52            read_state: DirectionState::new(read_acquirer, max_chunk),
53            write_state: DirectionState::new(write_acquirer, max_chunk),
54        }
55    }
56}
57
58impl<S> AsyncRead for GlobalRateLimitedConn<S>
59where
60    S: AsyncRead,
61{
62    fn poll_read(
63        self: Pin<&mut Self>,
64        cx: &mut Context<'_>,
65        buf: &mut [u8],
66    ) -> Poll<Result<usize, Error>> {
67        let this = self.project();
68        super::reader::poll_read_limited(this.inner, this.read_state, cx, buf)
69    }
70}
71
72impl<S> AsyncWrite for GlobalRateLimitedConn<S>
73where
74    S: AsyncWrite,
75{
76    fn poll_write(
77        self: Pin<&mut Self>,
78        cx: &mut Context<'_>,
79        buf: &[u8],
80    ) -> Poll<Result<usize, Error>> {
81        let this = self.project();
82        super::writer::poll_write_limited(this.inner, this.write_state, cx, buf)
83    }
84
85    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
86        self.project().inner.poll_flush(cx)
87    }
88
89    fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
90        self.project().inner.poll_close(cx)
91    }
92}
93
94/// Implement [`StreamOps`] and forward it to the inner stream.
95///
96/// Every operation is untouched, rate-limiting is not applied here.
97impl<S> StreamOps for GlobalRateLimitedConn<S>
98where
99    S: StreamOps,
100{
101    fn set_tcp_notsent_lowat(&self, notsent_lowat: u32) -> std::io::Result<()> {
102        self.inner.set_tcp_notsent_lowat(notsent_lowat)
103    }
104
105    fn new_handle(&self) -> Box<dyn StreamOps + Send + Unpin> {
106        self.inner.new_handle()
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 _, AsyncWriteExt as _, FutureExt as _, io::Cursor};
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    /// Build a conn over a [`Cursor`] with a 100 bytes length vector.
137    fn new_conn(
138        read_pool: &BandwidthPool,
139        write_pool: &BandwidthPool,
140    ) -> GlobalRateLimitedConn<Cursor<Vec<u8>>> {
141        GlobalRateLimitedConn::new(
142            Cursor::new(vec![1_u8; 100]),
143            read_pool.new_acquirer(),
144            write_pool.new_acquirer(),
145            MAX_CHUNK,
146        )
147    }
148
149    #[test]
150    fn basic_conn() {
151        let (read_pool, _rr) = BandwidthPool::new(100);
152        let (write_pool, _wr) = BandwidthPool::new(100);
153        let mut conn = new_conn(&read_pool, &write_pool);
154
155        // A read of 30 only spends from the read pool.
156        let mut buf = [0; 30];
157        assert_eq!(conn.read(&mut buf).now_or_never().unwrap().unwrap(), 30);
158        assert_eq!(read_pool.available(), 70);
159        assert_eq!(write_pool.available(), 100);
160
161        // A write of 40 only spends from the write pool.
162        assert_eq!(conn.write(&[1; 40]).now_or_never().unwrap().unwrap(), 40);
163        assert_eq!(read_pool.available(), 70);
164        assert_eq!(write_pool.available(), 60);
165    }
166
167    #[test]
168    fn write_no_read_block() {
169        let (read_pool, _rr) = BandwidthPool::new(50);
170        let (write_pool, _wr) = BandwidthPool::new(50);
171        let mut conn = new_conn(&read_pool, &write_pool);
172
173        // Drain the write pool.
174        assert_eq!(conn.write(&[1; 50]).now_or_never().unwrap().unwrap(), 50);
175
176        // A further write is Pending until a refill...
177        let mut write = conn.write(&[2; 50]);
178        assert!((&mut write).now_or_never().is_none());
179
180        // But a read is not blocked.
181        let mut buf = [0; 30];
182        assert_eq!(conn.read(&mut buf).now_or_never().unwrap().unwrap(), 30);
183        assert_eq!(read_pool.available(), 20);
184    }
185}