cloudiful-rate-limiter 0.1.0

Reusable async resource throttling with local and Valkey-backed backends.
Documentation
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);
    }
}