Skip to main content

tor_async_utils/global_rate_limit/
sink.rs

1//! A rate limited sink.
2//!
3//! A [`Sink`] wrapper that rate limits the items it forwards by acquiring bandwidth from
4//! a [`crate::bw_pool::BandwidthPool`] before each item is sent.
5//!
6//! # Example
7//!
8//! ```
9//! use futures::{SinkExt as _, channel::mpsc};
10//! use tor_async_utils::global_rate_limit::GlobalRateLimitedSink;
11//! use tor_async_utils::bw_pool::BandwidthPool;
12//!
13//! futures::executor::block_on(async {
14//!     let (pool, _refiller) = BandwidthPool::new(64 * 1024);
15//!     let (tx, _rx) = mpsc::channel::<Vec<u8>>(8);
16//!     let mut sink = GlobalRateLimitedSink::new(tx, pool.new_acquirer(), 512);
17//!
18//!     // The pool starts full so this is served from the fast path.
19//!     sink.send(vec![0; 512]).await.unwrap();
20//! });
21//! ```
22
23use 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/// A [`Sink`] wrapper that acquires bandwidth before forwarding each item.
32///
33/// Every item costs the same fixed number of `tokens` given to [`Self::new`], so this is
34/// intended for sinks whose items are of roughly the same size, such as Tor cells :).
35#[derive(Debug)]
36#[pin_project]
37pub struct GlobalRateLimitedSink<S> {
38    /// The underlying sink items are forwarded to.
39    #[pin]
40    inner: S,
41    /// Acquirer used to get a [`Permit`] from the pool for each item.
42    acquirer: BandwidthAcquirer,
43    /// The number of tokens each item costs, requested from the pool per item.
44    tokens: u64,
45    /// The permit for the next item acquired by [`Sink::poll_ready`].
46    ///
47    /// It is kept here across `poll_ready` calls so we never request a grant twice for
48    /// the same item. It is consumed by [`Sink::start_send`]. If the sink is dropped
49    /// a permit, it is refunded to the pool.
50    permit: Option<Permit>,
51}
52
53impl<S> GlobalRateLimitedSink<S> {
54    /// Construct a sink that spends `tokens` tokens from the `acquirer`'s pool for every
55    /// item it forwards.
56    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        // check the inner sink first so if it is backpressured, don't claim the tokens.
76        ready!(this.inner.poll_ready(cx)).map_err(GlobalRateLimitedError::Sink)?;
77
78        // If no permit, get one.
79        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        // Consume the permit that a successful poll_ready() acquired. We'll then claim
91        // all tokens as a permit is for the size of the item. We could use this.tokens
92        // but this protects us for the case where the number of tokens changed in
93        // between calls. Very unlikely but hey, safety first!
94        let mut permit = this
95            .permit
96            .take()
97            .ok_or(GlobalRateLimitedError::MissingPermit)?;
98        // TODO(relay): This can't work like this with partial permit. We need to check
99        // how many we got and claim that.
100        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    // @@ begin test lint list maintained by maint/add_warning @@
125    #![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)] // See arti#2571
137    //! <!-- @@ end test lint list maintained by maint/add_warning @@ -->
138
139    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        // The pool starts full: three sends go through the fast path...
153        for i in 0..3 {
154            sink.send(i).now_or_never().unwrap().unwrap();
155        }
156        // Three sends means 30 tokens each time (90 total) so 100 - 90 = 10 left.
157        assert_eq!(pool.available(), 10);
158
159        // Fourth send, pool is dry so we end up Pending.
160        let mut send4 = sink.send(4);
161        assert!((&mut send4).now_or_never().is_none());
162
163        // Make sure we got the three initial sends.
164        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        // Uses the full pool (30).
176        sink.send(1).now_or_never().unwrap().unwrap();
177
178        // Pool is empty. The send is Pending until refill.
179        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        // Make sure we got the two sends.
186        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        // Buffer of zero here means the capacity is 1 because 1 sender. That is from the
195        // channel() documentation.
196        let (tx, mut rx) = mpsc::channel::<u64>(0);
197        let mut sink = GlobalRateLimitedSink::new(tx, pool.new_acquirer(), 10);
198
199        // Use feed() so we don't flush. We just want to fill the channel.
200        sink.feed(1).now_or_never().unwrap().unwrap();
201        assert_eq!(pool.available(), 90);
202
203        // The channel is full. The feed() will be Pending before any tokens are claimed.
204        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        // Consuming the channel which should allow the next feed() to claim tokens.
210        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        // Drains the pool.
222        sink.send(1).now_or_never().unwrap().unwrap();
223        assert_eq!(pool.available(), 0);
224
225        // Without a refiller and no tokens left, the pool is closed and the sink errors.
226        drop(refiller);
227        assert!(matches!(
228            sink.send(2).now_or_never(),
229            Some(Err(GlobalRateLimitedError::Pool(_)))
230        ));
231    }
232}