tor_async_utils/global_rate_limit/
reader.rs1use 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#[derive(Debug)]
42#[pin_project]
43pub struct GlobalRateLimitedReader<R> {
44 #[pin]
46 inner: R,
47 state: DirectionState,
49}
50
51impl<R> GlobalRateLimitedReader<R> {
52 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
77pub(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 if buf.is_empty() {
87 return inner.poll_read(cx, buf);
88 }
89
90 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 Poll::Pending => Poll::Pending,
100 Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
102 Poll::Ready(Ok(read)) => {
104 DirectionState::commit(permit, read);
105 Poll::Ready(Ok(read))
106 }
107 }
108}
109
110#[cfg(test)]
111mod test {
112 #![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)] use super::*;
128
129 use futures::{AsyncReadExt as _, FutureExt as _};
130
131 use crate::bw_pool::BandwidthPool;
132
133 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 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 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 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 let mut buf = [0; 30];
181 assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 30);
182
183 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 let mut buf = [0; 10];
198 assert_eq!(reader.read(&mut buf).now_or_never().unwrap().unwrap(), 10);
199
200 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}