tor_async_utils/global_rate_limit/
conn.rs1use 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#[derive(Debug)]
25#[pin_project]
26pub struct GlobalRateLimitedConn<S> {
27 #[pin]
29 inner: S,
30 read_state: DirectionState,
32 write_state: DirectionState,
34}
35
36impl<S> GlobalRateLimitedConn<S> {
37 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
94impl<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 #![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 _, AsyncWriteExt as _, FutureExt as _, io::Cursor};
130
131 use crate::bw_pool::BandwidthPool;
132
133 const MAX_CHUNK: NonZero<usize> = NonZero::new(1024).unwrap();
135
136 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 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 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 assert_eq!(conn.write(&[1; 50]).now_or_never().unwrap().unwrap(), 50);
175
176 let mut write = conn.write(&[2; 50]);
178 assert!((&mut write).now_or_never().is_none());
179
180 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}