tor_async_utils/global_rate_limit/
sink.rs1use futures::Sink;
24use pin_project::pin_project;
25use std::pin::Pin;
26use std::task::{Context, Poll, ready};
27
28use crate::bw_pool::{BandwidthAcquirer, Permit};
29use crate::global_rate_limit::GlobalRateLimitedError;
30
31#[derive(Debug)]
36#[pin_project]
37pub struct GlobalRateLimitedSink<S> {
38 #[pin]
40 inner: S,
41 acquirer: BandwidthAcquirer,
43 tokens: u64,
45 permit: Option<Permit>,
51}
52
53impl<S> GlobalRateLimitedSink<S> {
54 pub fn new(inner: S, acquirer: BandwidthAcquirer, tokens: u64) -> Self {
57 Self {
58 inner,
59 acquirer,
60 tokens,
61 permit: None,
62 }
63 }
64}
65
66impl<S, Item> Sink<Item> for GlobalRateLimitedSink<S>
67where
68 S: Sink<Item>,
69{
70 type Error = GlobalRateLimitedError<S::Error>;
71
72 fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
73 let this = self.project();
74
75 ready!(this.inner.poll_ready(cx)).map_err(GlobalRateLimitedError::Sink)?;
77
78 if this.permit.is_none() {
80 let permit = ready!(this.acquirer.poll_acquire(cx, *this.tokens))?;
81 *this.permit = Some(permit);
82 }
83
84 Poll::Ready(Ok(()))
85 }
86
87 fn start_send(self: Pin<&mut Self>, item: Item) -> Result<(), Self::Error> {
88 let this = self.project();
89
90 let mut permit = this
95 .permit
96 .take()
97 .ok_or(GlobalRateLimitedError::MissingPermit)?;
98 permit.claim_all();
101
102 this.inner
103 .start_send(item)
104 .map_err(GlobalRateLimitedError::Sink)
105 }
106
107 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
108 self.project()
109 .inner
110 .poll_flush(cx)
111 .map_err(GlobalRateLimitedError::Sink)
112 }
113
114 fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
115 self.project()
116 .inner
117 .poll_close(cx)
118 .map_err(GlobalRateLimitedError::Sink)
119 }
120}
121
122#[cfg(test)]
123mod test {
124 #![allow(clippy::bool_assert_comparison)]
126 #![allow(clippy::clone_on_copy)]
127 #![allow(clippy::dbg_macro)]
128 #![allow(clippy::mixed_attributes_style)]
129 #![allow(clippy::print_stderr)]
130 #![allow(clippy::print_stdout)]
131 #![allow(clippy::single_char_pattern)]
132 #![allow(clippy::unwrap_used)]
133 #![allow(clippy::unchecked_time_subtraction)]
134 #![allow(clippy::useless_vec)]
135 #![allow(clippy::needless_pass_by_value)]
136 #![allow(clippy::string_slice)] use super::*;
140
141 use futures::channel::mpsc;
142 use futures::{FutureExt as _, SinkExt as _, StreamExt as _};
143
144 use crate::bw_pool::BandwidthPool;
145
146 #[test]
147 fn fast_path() {
148 let (pool, _refiller) = BandwidthPool::new(100);
149 let (tx, mut rx) = mpsc::channel::<u64>(4);
150 let mut sink = GlobalRateLimitedSink::new(tx, pool.new_acquirer(), 30);
151
152 for i in 0..3 {
154 sink.send(i).now_or_never().unwrap().unwrap();
155 }
156 assert_eq!(pool.available(), 10);
158
159 let mut send4 = sink.send(4);
161 assert!((&mut send4).now_or_never().is_none());
162
163 for i in 0..3 {
165 assert_eq!(rx.next().now_or_never().unwrap(), Some(i));
166 }
167 }
168
169 #[test]
170 fn fast_path_refill() {
171 let (pool, mut refiller) = BandwidthPool::new(30);
172 let (tx, mut rx) = mpsc::channel::<u64>(4);
173 let mut sink = GlobalRateLimitedSink::new(tx, pool.new_acquirer(), 30);
174
175 sink.send(1).now_or_never().unwrap().unwrap();
177
178 let mut send2 = sink.send(2);
180 assert!((&mut send2).now_or_never().is_none());
181 assert_eq!(refiller.refill_and_serve(30), None);
182 assert!(matches!((&mut send2).now_or_never(), Some(Ok(()))));
183 drop(send2);
184
185 for i in 0..2 {
187 assert_eq!(rx.next().now_or_never().unwrap(), Some(i + 1));
188 }
189 }
190
191 #[test]
192 fn backpressure() {
193 let (pool, _refiller) = BandwidthPool::new(100);
194 let (tx, mut rx) = mpsc::channel::<u64>(0);
197 let mut sink = GlobalRateLimitedSink::new(tx, pool.new_acquirer(), 10);
198
199 sink.feed(1).now_or_never().unwrap().unwrap();
201 assert_eq!(pool.available(), 90);
202
203 let mut feed2 = sink.feed(2);
205 assert!((&mut feed2).now_or_never().is_none());
206 drop(feed2);
207 assert_eq!(pool.available(), 90);
208
209 assert_eq!(rx.next().now_or_never().unwrap(), Some(1));
211 sink.feed(2).now_or_never().unwrap().unwrap();
212 assert_eq!(pool.available(), 80);
213 }
214
215 #[test]
216 fn pool_closed() {
217 let (pool, refiller) = BandwidthPool::new(10);
218 let (tx, _rx) = mpsc::channel::<u64>(4);
219 let mut sink = GlobalRateLimitedSink::new(tx, pool.new_acquirer(), 10);
220
221 sink.send(1).now_or_never().unwrap().unwrap();
223 assert_eq!(pool.available(), 0);
224
225 drop(refiller);
227 assert!(matches!(
228 sink.send(2).now_or_never(),
229 Some(Err(GlobalRateLimitedError::Pool(_)))
230 ));
231 }
232}