use std::time::{Duration, Instant};
#[derive(Debug)]
pub struct RateLimiter {
capacity_per_minute: u64,
tokens_per_sec: f64,
tokens: f64,
last_check: Instant,
}
impl RateLimiter {
pub fn new_per_minute(requests_num: usize) -> Self {
let tokens_per_sec = requests_num as f64 / 60.0;
RateLimiter {
capacity_per_minute: requests_num as u64,
tokens_per_sec,
tokens: requests_num as f64, last_check: Instant::now(),
}
}
pub fn try_consume(&mut self, tokens: f64) -> Result<(), RateLimitError> {
if tokens > self.capacity_per_minute as f64 {
return Err(RateLimitError::AlwaysOverBudget(
"request larger than rate limiter capacity, please try to split your request",
));
}
let now = Instant::now();
let elapsed = now.duration_since(self.last_check);
self.last_check = now;
self.tokens += self.tokens_per_sec * elapsed.as_secs_f64();
if self.tokens > self.capacity_per_minute as f64 {
self.tokens = self.capacity_per_minute as f64;
}
if self.tokens >= tokens {
self.tokens -= tokens; Ok(()) } else {
let missing_tokens = tokens - self.tokens;
let retry_after = Duration::from_secs_f64(missing_tokens / self.tokens_per_sec);
debug_assert!(retry_after > Duration::from_secs(0));
let retry_error = RetryError {
tokens_available: self.tokens,
retry_after,
};
Err(RateLimitError::Retry(retry_error))
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct RetryError {
pub tokens_available: f64,
pub retry_after: Duration,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum RateLimitError {
AlwaysOverBudget(&'static str),
Retry(RetryError),
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_eq_floats(a: f64, b: f64, tolerance: f64) {
assert!(
(a - b).abs() < tolerance,
"assertion failed: `(left == right)` (left: `{a}`, right: `{b}`, tolerance: `{tolerance}`)",
);
}
#[test]
fn test_rate_one_per_minute() {
let mut limiter = RateLimiter::new_per_minute(1);
assert_eq!(limiter.capacity_per_minute, 1);
assert_eq_floats(limiter.tokens_per_sec, 0.016, 0.001);
assert_eq!(limiter.tokens, 1.0);
assert_eq!(limiter.try_consume(1.0), Ok(()));
assert_eq!(limiter.tokens, 0.0);
assert!(limiter.try_consume(1.0).is_err());
}
#[test]
fn test_rate_more_per_minute() {
let mut limiter = RateLimiter::new_per_minute(600);
assert_eq!(limiter.capacity_per_minute, 600);
assert_eq!(limiter.tokens_per_sec, 10.0);
assert_eq!(limiter.tokens, 600.0);
assert_eq!(limiter.try_consume(1.0), Ok(()));
assert_eq!(limiter.tokens, 599.0);
assert_eq!(limiter.try_consume(10.0), Ok(()));
assert!((589.0..=590.0).contains(&limiter.tokens));
}
#[test]
fn test_rate_huge_request() {
let mut limiter = RateLimiter::new_per_minute(100);
assert!(limiter.try_consume(99999.0).is_err());
}
}