tor_async_utils/global_rate_limit/
writer.rs1use 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#[derive(Debug)]
40#[pin_project]
41pub struct GlobalRateLimitedWriter<W> {
42 #[pin]
44 inner: W,
45 state: DirectionState,
47}
48
49impl<W> GlobalRateLimitedWriter<W> {
50 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
85pub(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 if buf.is_empty() {
95 return inner.poll_write(cx, buf);
96 }
97
98 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 Poll::Pending => Poll::Pending,
106 Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
108 Poll::Ready(Ok(written)) => {
110 DirectionState::commit(permit, written);
111 Poll::Ready(Ok(written))
112 }
113 }
114}
115
116#[cfg(test)]
117mod test {
118 #![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)] use super::*;
134
135 use futures::{AsyncWriteExt as _, FutureExt as _};
136
137 use crate::bw_pool::BandwidthPool;
138
139 const MAX_CHUNK: NonZero<usize> = NonZero::new(1024).unwrap();
141
142 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 assert!(writer.write(&[0; 30]).now_or_never().is_none());
169
170 assert_eq!(pool.available(), 100);
173
174 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 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 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 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 assert_eq!(writer.write(&[0; 30]).now_or_never().unwrap().unwrap(), 30);
217
218 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 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 assert_eq!(writer.write(&[0; 10]).now_or_never().unwrap().unwrap(), 10);
233
234 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}