Skip to main content

hey_sdk/resilience/
rate_limit.rs

1use std::sync::{Mutex, PoisonError};
2use std::time::{Duration, Instant};
3
4use super::Clock;
5
6/// The budget a client holds itself to, and whether it takes HEY's word over its own.
7///
8/// A rate or a burst that is not positive reads as "leave it at the default", the same
9/// normalising Go does.
10#[derive(Debug, Clone)]
11pub struct RateLimitConfig {
12    /// How fast the bucket refills.
13    pub requests_per_second: f64,
14    /// How many requests can go out at once from a full bucket.
15    pub burst_size: u32,
16    /// Whether a `Retry-After` from HEY holds the limiter back until it has passed.
17    pub respect_retry_after: bool,
18    /// Where the limiter reads the time; see [`Clock`].
19    pub clock: Clock,
20}
21
22impl Default for RateLimitConfig {
23    fn default() -> RateLimitConfig {
24        RateLimitConfig {
25            requests_per_second: 50.0,
26            burst_size: 10,
27            respect_retry_after: true,
28            clock: Clock::default(),
29        }
30    }
31}
32
33/// A token bucket, and whatever wait HEY last asked for.
34///
35/// The bucket starts full at [`RateLimitConfig::burst_size`] and refills at
36/// [`RateLimitConfig::requests_per_second`]; every call spends a token. While a
37/// `Retry-After` is in force nothing goes out at all, however many tokens have piled up.
38pub struct RateLimiter {
39    requests_per_second: f64,
40    burst_size: f64,
41    respect_retry_after: bool,
42    clock: Clock,
43    inner: Mutex<Bucket>,
44}
45
46struct Bucket {
47    tokens: f64,
48    last_refill: Instant,
49    retry_after_until: Option<Instant>,
50}
51
52impl RateLimiter {
53    /// A limiter at `config`, its bucket full.
54    pub fn new(config: RateLimitConfig) -> RateLimiter {
55        let defaults = RateLimitConfig::default();
56        let burst_size = match config.burst_size {
57            0 => defaults.burst_size,
58            burst_size => burst_size,
59        };
60        let bucket = Bucket {
61            tokens: f64::from(burst_size),
62            last_refill: config.clock.now(),
63            retry_after_until: None,
64        };
65        RateLimiter {
66            requests_per_second: if config.requests_per_second > 0.0 {
67                config.requests_per_second
68            } else {
69                defaults.requests_per_second
70            },
71            burst_size: f64::from(burst_size),
72            respect_retry_after: config.respect_retry_after,
73            clock: config.clock,
74            inner: Mutex::new(bucket),
75        }
76    }
77
78    /// Whether a call may go out now, spending a token if it may.
79    pub fn allow(&self) -> bool {
80        let mut bucket = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
81        if self.held_back(&mut bucket) {
82            false
83        } else {
84            self.refill(&mut bucket);
85            if bucket.tokens >= 1.0 {
86                bucket.tokens -= 1.0;
87                true
88            } else {
89                false
90            }
91        }
92    }
93
94    /// How long a call would have to wait for its token: `Some(Duration::ZERO)` to go now,
95    /// and `None` when the limiter will not have one within the second — or when HEY has
96    /// asked for a wait, which the caller should sit out rather than shorten.
97    ///
98    /// A reservation spends its token up front, as Go's does, so a caller that reserves and
99    /// then walks away leaves the bucket short.
100    pub fn reserve(&self) -> Option<Duration> {
101        let mut bucket = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
102        if self.held_back(&mut bucket) {
103            return None;
104        }
105
106        self.refill(&mut bucket);
107        if bucket.tokens >= 1.0 {
108            bucket.tokens -= 1.0;
109            return Some(Duration::ZERO);
110        }
111
112        let wait = Duration::from_secs_f64((1.0 - bucket.tokens) / self.requests_per_second);
113        if wait > Duration::from_secs(1) {
114            None
115        } else {
116            bucket.tokens -= 1.0;
117            Some(wait)
118        }
119    }
120
121    /// Holds every call back until the given moment. A wait already in force is only ever
122    /// extended, never cut short, so the longest thing HEY asked for is what is honoured.
123    pub fn set_retry_after(&self, until: Instant) {
124        if self.respect_retry_after {
125            let mut bucket = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
126            if bucket
127                .retry_after_until
128                .is_none_or(|current| until > current)
129            {
130                bucket.retry_after_until = Some(until);
131            }
132        }
133    }
134
135    /// [`RateLimiter::set_retry_after`], measured from now.
136    pub fn set_retry_after_in(&self, wait: Duration) {
137        self.set_retry_after(self.clock.now() + wait);
138    }
139
140    /// How much of the wait HEY asked for is left, and zero when it asked for none.
141    pub fn retry_after_remaining(&self) -> Duration {
142        match self
143            .inner
144            .lock()
145            .unwrap_or_else(PoisonError::into_inner)
146            .retry_after_until
147        {
148            Some(until) => until.saturating_duration_since(self.clock.now()),
149            None => Duration::ZERO,
150        }
151    }
152
153    /// How many calls the bucket would let through right now.
154    pub fn tokens(&self) -> f64 {
155        let mut bucket = self.inner.lock().unwrap_or_else(PoisonError::into_inner);
156        self.refill(&mut bucket);
157        bucket.tokens
158    }
159
160    /// Whether HEY's wait is still in force, forgetting one that has passed on the way.
161    fn held_back(&self, bucket: &mut Bucket) -> bool {
162        match bucket.retry_after_until {
163            Some(until) if self.respect_retry_after => {
164                if self.clock.now() < until {
165                    true
166                } else {
167                    bucket.retry_after_until = None;
168                    false
169                }
170            }
171            _ => false,
172        }
173    }
174
175    fn refill(&self, bucket: &mut Bucket) {
176        let now = self.clock.now();
177        let elapsed = now.saturating_duration_since(bucket.last_refill);
178        bucket.last_refill = now;
179        bucket.tokens =
180            (bucket.tokens + elapsed.as_secs_f64() * self.requests_per_second).min(self.burst_size);
181    }
182}
183
184#[cfg(test)]
185mod tests {
186    use super::super::{advance, test_clock};
187    use super::*;
188
189    #[test]
190    fn a_burst_goes_out_at_once_and_the_next_call_waits_for_the_refill() {
191        let (clock, now) = test_clock();
192        let limiter = RateLimiter::new(RateLimitConfig {
193            requests_per_second: 100.0,
194            burst_size: 5,
195            clock,
196            ..RateLimitConfig::default()
197        });
198
199        for spent in 1..=5 {
200            assert!(limiter.allow(), "call {spent} of the burst");
201        }
202        assert!(!limiter.allow());
203
204        advance(&now, Duration::from_millis(20));
205
206        assert!(limiter.allow());
207    }
208
209    #[test]
210    fn a_wait_hey_asked_for_holds_everything_back_until_it_has_passed() {
211        let (clock, now) = test_clock();
212        let limiter = RateLimiter::new(RateLimitConfig {
213            requests_per_second: 1000.0,
214            burst_size: 100,
215            respect_retry_after: true,
216            clock,
217        });
218
219        limiter.set_retry_after_in(Duration::from_secs(1));
220
221        assert!(!limiter.allow());
222        assert!(limiter.retry_after_remaining() > Duration::ZERO);
223
224        advance(&now, Duration::from_secs(2));
225
226        assert!(limiter.allow());
227        assert_eq!(Duration::ZERO, limiter.retry_after_remaining());
228    }
229
230    #[test]
231    fn a_limiter_told_not_to_respect_the_wait_carries_on() {
232        let (clock, _now) = test_clock();
233        let limiter = RateLimiter::new(RateLimitConfig {
234            respect_retry_after: false,
235            clock,
236            ..RateLimitConfig::default()
237        });
238
239        limiter.set_retry_after_in(Duration::from_secs(60));
240
241        assert!(limiter.allow());
242        assert_eq!(Duration::ZERO, limiter.retry_after_remaining());
243    }
244
245    #[test]
246    fn a_longer_wait_replaces_a_shorter_one_and_a_shorter_one_is_ignored() {
247        let (clock, _now) = test_clock();
248        let limiter = RateLimiter::new(RateLimitConfig {
249            clock,
250            ..RateLimitConfig::default()
251        });
252
253        limiter.set_retry_after_in(Duration::from_secs(30));
254        limiter.set_retry_after_in(Duration::from_secs(5));
255        assert_eq!(Duration::from_secs(30), limiter.retry_after_remaining());
256
257        limiter.set_retry_after_in(Duration::from_secs(60));
258        assert_eq!(Duration::from_secs(60), limiter.retry_after_remaining());
259    }
260
261    #[test]
262    #[allow(clippy::float_cmp)] // whole tokens, and the clock has not moved to add a fraction
263    fn a_new_bucket_is_full_and_every_call_spends_from_it() {
264        let (clock, _now) = test_clock();
265        let limiter = RateLimiter::new(RateLimitConfig {
266            requests_per_second: 100.0,
267            burst_size: 10,
268            clock,
269            ..RateLimitConfig::default()
270        });
271
272        assert_eq!(10.0, limiter.tokens());
273
274        limiter.allow();
275
276        assert_eq!(9.0, limiter.tokens());
277    }
278
279    #[test]
280    fn a_reservation_answers_the_wait_its_token_costs() {
281        let (clock, _now) = test_clock();
282        let limiter = RateLimiter::new(RateLimitConfig {
283            requests_per_second: 10.0,
284            burst_size: 1,
285            clock,
286            ..RateLimitConfig::default()
287        });
288
289        assert_eq!(Some(Duration::ZERO), limiter.reserve());
290        assert_eq!(Some(Duration::from_millis(100)), limiter.reserve());
291        assert_eq!(Some(Duration::from_millis(200)), limiter.reserve());
292    }
293
294    #[test]
295    fn a_reservation_beyond_a_second_out_is_refused() {
296        let (clock, _now) = test_clock();
297        let limiter = RateLimiter::new(RateLimitConfig {
298            requests_per_second: 0.5,
299            burst_size: 1,
300            clock,
301            ..RateLimitConfig::default()
302        });
303
304        assert_eq!(Some(Duration::ZERO), limiter.reserve());
305
306        assert_eq!(None, limiter.reserve());
307    }
308
309    #[test]
310    fn a_reservation_is_refused_while_hey_has_asked_for_a_wait() {
311        let (clock, _now) = test_clock();
312        let limiter = RateLimiter::new(RateLimitConfig {
313            clock,
314            ..RateLimitConfig::default()
315        });
316
317        limiter.set_retry_after_in(Duration::from_secs(30));
318
319        assert_eq!(None, limiter.reserve());
320    }
321
322    #[test]
323    #[allow(clippy::float_cmp)] // whole tokens, and the clock has not moved to add a fraction
324    fn a_config_of_zeroes_falls_back_to_the_defaults() {
325        let (clock, _now) = test_clock();
326        let limiter = RateLimiter::new(RateLimitConfig {
327            requests_per_second: 0.0,
328            burst_size: 0,
329            respect_retry_after: true,
330            clock,
331        });
332
333        assert_eq!(10.0, limiter.tokens());
334    }
335}