1use std::sync::Mutex;
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
81 .inner
82 .lock()
83 .unwrap_or_else(std::sync::PoisonError::into_inner);
84 if self.held_back(&mut bucket) {
85 false
86 } else {
87 self.refill(&mut bucket);
88 if bucket.tokens >= 1.0 {
89 bucket.tokens -= 1.0;
90 true
91 } else {
92 false
93 }
94 }
95 }
96
97 pub fn reserve(&self) -> Option<Duration> {
104 let mut bucket = self
105 .inner
106 .lock()
107 .unwrap_or_else(std::sync::PoisonError::into_inner);
108 if self.held_back(&mut bucket) {
109 return None;
110 }
111
112 self.refill(&mut bucket);
113 if bucket.tokens >= 1.0 {
114 bucket.tokens -= 1.0;
115 return Some(Duration::ZERO);
116 }
117
118 let wait = Duration::from_secs_f64((1.0 - bucket.tokens) / self.requests_per_second);
119 if wait > Duration::from_secs(1) {
120 None
121 } else {
122 bucket.tokens -= 1.0;
123 Some(wait)
124 }
125 }
126
127 pub fn set_retry_after(&self, until: Instant) {
130 if self.respect_retry_after {
131 let mut bucket = self
132 .inner
133 .lock()
134 .unwrap_or_else(std::sync::PoisonError::into_inner);
135 if bucket
136 .retry_after_until
137 .is_none_or(|current| until > current)
138 {
139 bucket.retry_after_until = Some(until);
140 }
141 }
142 }
143
144 pub fn set_retry_after_in(&self, wait: Duration) {
146 let until = self
148 .clock
149 .now()
150 .checked_add(wait)
151 .unwrap_or_else(|| self.clock.now() + Duration::from_secs(u32::MAX.into()));
152 self.set_retry_after(until);
153 }
154
155 pub fn retry_after_remaining(&self) -> Duration {
157 match self
158 .inner
159 .lock()
160 .unwrap_or_else(std::sync::PoisonError::into_inner)
161 .retry_after_until
162 {
163 Some(until) => until.saturating_duration_since(self.clock.now()),
164 None => Duration::ZERO,
165 }
166 }
167
168 pub fn tokens(&self) -> f64 {
170 let mut bucket = self
171 .inner
172 .lock()
173 .unwrap_or_else(std::sync::PoisonError::into_inner);
174 self.refill(&mut bucket);
175 bucket.tokens
176 }
177
178 fn held_back(&self, bucket: &mut Bucket) -> bool {
180 match bucket.retry_after_until {
181 Some(until) if self.respect_retry_after => {
182 if self.clock.now() < until {
183 true
184 } else {
185 bucket.retry_after_until = None;
186 false
187 }
188 }
189 _ => false,
190 }
191 }
192
193 fn refill(&self, bucket: &mut Bucket) {
194 let now = self.clock.now();
195 let elapsed = now.saturating_duration_since(bucket.last_refill);
196 bucket.last_refill = now;
197 bucket.tokens =
198 (bucket.tokens + elapsed.as_secs_f64() * self.requests_per_second).min(self.burst_size);
199 }
200}
201
202#[cfg(test)]
203#[allow(clippy::unwrap_used, clippy::float_cmp)]
204mod tests {
205 use super::super::{advance, test_clock};
206 use super::*;
207
208 #[test]
209 fn a_burst_goes_out_at_once_and_the_next_call_waits_for_the_refill() {
210 let (clock, now) = test_clock();
211 let limiter = RateLimiter::new(RateLimitConfig {
212 requests_per_second: 100.0,
213 burst_size: 5,
214 clock,
215 ..RateLimitConfig::default()
216 });
217
218 for spent in 1..=5 {
219 assert!(limiter.allow(), "call {spent} of the burst");
220 }
221 assert!(!limiter.allow());
222
223 advance(&now, Duration::from_millis(20));
224
225 assert!(limiter.allow());
226 }
227
228 #[test]
229 fn a_wait_fizzy_asked_for_holds_everything_back_until_it_has_passed() {
230 let (clock, now) = test_clock();
231 let limiter = RateLimiter::new(RateLimitConfig {
232 requests_per_second: 1000.0,
233 burst_size: 100,
234 respect_retry_after: true,
235 clock,
236 });
237
238 limiter.set_retry_after_in(Duration::from_secs(1));
239
240 assert!(!limiter.allow());
241 assert!(limiter.retry_after_remaining() > Duration::ZERO);
242
243 advance(&now, Duration::from_secs(2));
244
245 assert!(limiter.allow());
246 assert_eq!(Duration::ZERO, limiter.retry_after_remaining());
247 }
248
249 #[test]
250 fn a_limiter_told_not_to_respect_the_wait_carries_on() {
251 let (clock, _now) = test_clock();
252 let limiter = RateLimiter::new(RateLimitConfig {
253 respect_retry_after: false,
254 clock,
255 ..RateLimitConfig::default()
256 });
257
258 limiter.set_retry_after_in(Duration::from_secs(60));
259
260 assert!(limiter.allow());
261 assert_eq!(Duration::ZERO, limiter.retry_after_remaining());
262 }
263
264 #[test]
265 fn a_longer_wait_replaces_a_shorter_one_and_a_shorter_one_is_ignored() {
266 let (clock, _now) = test_clock();
267 let limiter = RateLimiter::new(RateLimitConfig {
268 clock,
269 ..RateLimitConfig::default()
270 });
271
272 limiter.set_retry_after_in(Duration::from_secs(30));
273 limiter.set_retry_after_in(Duration::from_secs(5));
274 assert_eq!(Duration::from_secs(30), limiter.retry_after_remaining());
275
276 limiter.set_retry_after_in(Duration::from_secs(60));
277 assert_eq!(Duration::from_secs(60), limiter.retry_after_remaining());
278 }
279
280 #[test]
281 fn a_new_bucket_is_full_and_every_call_spends_from_it() {
282 let (clock, _now) = test_clock();
283 let limiter = RateLimiter::new(RateLimitConfig {
284 requests_per_second: 100.0,
285 burst_size: 10,
286 clock,
287 ..RateLimitConfig::default()
288 });
289
290 assert_eq!(10.0, limiter.tokens());
291
292 limiter.allow();
293
294 assert_eq!(9.0, limiter.tokens());
295 }
296
297 #[test]
298 fn a_reservation_answers_the_wait_its_token_costs() {
299 let (clock, _now) = test_clock();
300 let limiter = RateLimiter::new(RateLimitConfig {
301 requests_per_second: 10.0,
302 burst_size: 1,
303 clock,
304 ..RateLimitConfig::default()
305 });
306
307 assert_eq!(Some(Duration::ZERO), limiter.reserve());
308 assert_eq!(Some(Duration::from_millis(100)), limiter.reserve());
309 assert_eq!(Some(Duration::from_millis(200)), limiter.reserve());
310 }
311
312 #[test]
313 fn a_reservation_beyond_a_second_out_is_refused() {
314 let (clock, _now) = test_clock();
315 let limiter = RateLimiter::new(RateLimitConfig {
316 requests_per_second: 0.5,
317 burst_size: 1,
318 clock,
319 ..RateLimitConfig::default()
320 });
321
322 assert_eq!(Some(Duration::ZERO), limiter.reserve());
323
324 assert_eq!(None, limiter.reserve());
325 }
326
327 #[test]
328 fn a_reservation_is_refused_while_fizzy_has_asked_for_a_wait() {
329 let (clock, _now) = test_clock();
330 let limiter = RateLimiter::new(RateLimitConfig {
331 clock,
332 ..RateLimitConfig::default()
333 });
334
335 limiter.set_retry_after_in(Duration::from_secs(30));
336
337 assert_eq!(None, limiter.reserve());
338 }
339
340 #[test]
341 fn a_config_of_zeroes_falls_back_to_the_defaults() {
342 let (clock, _now) = test_clock();
343 let limiter = RateLimiter::new(RateLimitConfig {
344 requests_per_second: 0.0,
345 burst_size: 0,
346 respect_retry_after: true,
347 clock,
348 });
349
350 assert_eq!(10.0, limiter.tokens());
351 }
352}