1use std::sync::{Mutex, PoisonError};
2use std::time::{Duration, Instant};
3
4use super::Clock;
5
6#[derive(Debug, Clone)]
11pub struct RateLimitConfig {
12 pub requests_per_second: f64,
14 pub burst_size: u32,
16 pub respect_retry_after: bool,
18 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
33pub 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 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 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 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 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 pub fn set_retry_after_in(&self, wait: Duration) {
137 self.set_retry_after(self.clock.now() + wait);
138 }
139
140 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 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 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)] 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)] 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}