use crate::{
AcquireResult, PeekResult, RateLimitError, RateLimitPolicy, RateLimiter, TryAcquireResult,
};
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
#[derive(Debug, Clone)]
pub struct LocalRateLimiter {
interval: Duration,
state: Arc<Mutex<HashMap<String, Instant>>>,
}
impl LocalRateLimiter {
pub fn new(policy: RateLimitPolicy) -> Result<Self, RateLimitError> {
let interval = policy.interval()?;
Ok(Self {
interval,
state: Arc::new(Mutex::new(HashMap::new())),
})
}
async fn reserve_slot(&self, key: &str) -> (Duration, Instant) {
let now = Instant::now();
let mut state = self.state.lock().await;
let next_allowed = state.get(key).copied().unwrap_or(now);
let wait_duration = next_allowed.saturating_duration_since(now);
let acquired_at = now + wait_duration;
state.insert(key.to_owned(), acquired_at + self.interval);
(wait_duration, acquired_at)
}
async fn try_reserve_slot(&self, key: &str) -> TryAcquireResult {
let now = Instant::now();
let mut state = self.state.lock().await;
let next_allowed = state.get(key).copied().unwrap_or(now);
let wait_duration = next_allowed.saturating_duration_since(now);
if !wait_duration.is_zero() {
return TryAcquireResult::Limited { wait_duration };
}
state.insert(key.to_owned(), now + self.interval);
TryAcquireResult::Acquired
}
async fn inspect(&self, key: &str) -> PeekResult {
let now = Instant::now();
let state = self.state.lock().await;
let wait_duration = state
.get(key)
.copied()
.unwrap_or(now)
.saturating_duration_since(now);
PeekResult {
allowed: wait_duration.is_zero(),
wait_duration,
}
}
}
#[async_trait]
impl RateLimiter for LocalRateLimiter {
async fn acquire(&self, key: &str) -> Result<AcquireResult, RateLimitError> {
let (waited, _) = self.reserve_slot(key).await;
if !waited.is_zero() {
tokio::time::sleep(waited).await;
}
Ok(AcquireResult {
key: key.to_owned(),
waited,
})
}
async fn try_acquire(&self, key: &str) -> Result<TryAcquireResult, RateLimitError> {
Ok(self.try_reserve_slot(key).await)
}
async fn peek(&self, key: &str) -> Result<PeekResult, RateLimitError> {
Ok(self.inspect(key).await)
}
}
#[cfg(test)]
mod tests {
use super::LocalRateLimiter;
use crate::{RateLimitPolicy, RateLimiter, TryAcquireResult};
use std::time::Duration;
fn assert_duration_at_least(actual: Duration, expected: Duration) {
assert!(actual >= expected.saturating_sub(Duration::from_millis(1)));
assert!(actual <= expected);
}
#[tokio::test(start_paused = true)]
async fn same_key_waits_on_second_acquire() {
let limiter =
LocalRateLimiter::new(RateLimitPolicy::min_interval(Duration::from_secs(3)).unwrap())
.unwrap();
let first = limiter.acquire("same").await.unwrap();
assert_eq!(first.waited, Duration::ZERO);
let limiter_clone = limiter.clone();
let task = tokio::spawn(async move { limiter_clone.acquire("same").await.unwrap() });
tokio::task::yield_now().await;
assert!(!task.is_finished());
tokio::time::advance(Duration::from_secs(3)).await;
let second = task.await.unwrap();
assert_duration_at_least(second.waited, Duration::from_secs(3));
}
#[tokio::test(start_paused = true)]
async fn different_keys_do_not_block_each_other() {
let limiter =
LocalRateLimiter::new(RateLimitPolicy::min_interval(Duration::from_secs(2)).unwrap())
.unwrap();
limiter.acquire("a").await.unwrap();
let second = limiter.acquire("b").await.unwrap();
assert_eq!(second.waited, Duration::ZERO);
}
#[tokio::test(start_paused = true)]
async fn try_acquire_reports_wait_duration() {
let limiter =
LocalRateLimiter::new(RateLimitPolicy::min_interval(Duration::from_secs(5)).unwrap())
.unwrap();
assert!(matches!(
limiter.try_acquire("same").await.unwrap(),
TryAcquireResult::Acquired
));
match limiter.try_acquire("same").await.unwrap() {
TryAcquireResult::Limited { wait_duration } => {
assert_duration_at_least(wait_duration, Duration::from_secs(5));
}
TryAcquireResult::Acquired => panic!("expected limited result"),
}
}
#[tokio::test(start_paused = true)]
async fn peek_does_not_consume_slot() {
let limiter =
LocalRateLimiter::new(RateLimitPolicy::min_interval(Duration::from_secs(4)).unwrap())
.unwrap();
let peek = limiter.peek("diag").await.unwrap();
assert!(peek.allowed);
assert_eq!(peek.wait_duration, Duration::ZERO);
let acquire = limiter.acquire("diag").await.unwrap();
assert_eq!(acquire.waited, Duration::ZERO);
}
}