kegani 0.1.1

A developer-friendly, ergonomic, production-ready Rust web framework
Documentation
//! Rate limiting middleware
//!
//! Provides request rate limiting using token bucket algorithm.

use actix_web::{
    dev::{Service, ServiceRequest, ServiceResponse, Transform},
    Error,
};
use std::future::{ready, Ready};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Instant;
use tokio::sync::Mutex;

/// Rate limit configuration
#[derive(Debug, Clone)]
pub struct RateLimitConfig {
    /// Requests per second
    pub requests_per_second: u32,
    /// Burst size
    pub burst_size: u32,
}

impl Default for RateLimitConfig {
    fn default() -> Self {
        Self::new(100)
    }
}

impl RateLimitConfig {
    /// Create a new config with the given requests per second
    pub fn new(requests_per_second: u32) -> Self {
        Self {
            requests_per_second,
            burst_size: requests_per_second * 2,
        }
    }
}

/// Token bucket state
struct TokenBucket {
    tokens: f64,
    last_update: Instant,
}

impl TokenBucket {
    fn new(max_tokens: f64) -> Self {
        Self {
            tokens: max_tokens,
            last_update: Instant::now(),
        }
    }

    fn try_consume(&mut self, tokens: f64, refill_rate: f64) -> bool {
        let now = Instant::now();
        let elapsed = now.duration_since(self.last_update).as_secs_f64();

        // Refill tokens
        self.tokens = (self.tokens + elapsed * refill_rate).min(self.tokens.max(tokens));
        self.last_update = now;

        if self.tokens >= tokens {
            self.tokens -= tokens;
            true
        } else {
            false
        }
    }
}

/// Rate limit middleware
pub struct RateLimit {
    config: RateLimitConfig,
    buckets: Arc<Mutex<Vec<TokenBucket>>>,
}

impl RateLimit {
    /// Create a new RateLimit middleware
    pub fn new(config: RateLimitConfig) -> Self {
        Self {
            buckets: Arc::new(Mutex::new(Vec::new())),
            config,
        }
    }

    /// Create with requests per second
    pub fn per_second(requests_per_second: u32) -> Self {
        Self::new(RateLimitConfig::new(requests_per_second))
    }
}

impl Default for RateLimit {
    fn default() -> Self {
        Self::new(RateLimitConfig::default())
    }
}

impl<S, B> Transform<S, ServiceRequest> for RateLimit
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static + Clone,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = Error;
    type Transform = RateLimitMiddleware<S>;
    type InitError = ();
    type Future = Ready<Result<Self::Transform, Self::InitError>>;

    fn new_transform(&self, service: S) -> Self::Future {
        ready(Ok(RateLimitMiddleware {
            service,
            config: self.config.clone(),
            buckets: self.buckets.clone(),
        }))
    }
}

pub struct RateLimitMiddleware<S> {
    service: S,
    config: RateLimitConfig,
    buckets: Arc<Mutex<Vec<TokenBucket>>>,
}

impl<S, B> Service<ServiceRequest> for RateLimitMiddleware<S>
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = Error;
    type Future = Pin<Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>>>>;

    fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
        self.service.poll_ready(cx)
    }

    fn call(&self, req: ServiceRequest) -> Self::Future {
        let client_ip = req
            .peer_addr()
            .map(|addr| addr.ip().to_string())
            .unwrap_or_else(|| "default".to_string());

        let hash = client_ip.bytes().fold(0u64, |acc, b| acc.wrapping_add(b as u64)) as usize;

        let config = self.config.clone();
        let buckets = self.buckets.clone();
        let fut = self.service.call(req);

        Box::pin(async move {
            let mut buckets = buckets.lock().await;

            // Ensure bucket exists
            while buckets.len() <= hash % 100 {
                buckets.push(TokenBucket::new(config.burst_size as f64));
            }

            let bucket = &mut buckets[hash % 100];
            let refill_rate = config.requests_per_second as f64;

            if bucket.try_consume(1.0, refill_rate) {
                fut.await
            } else {
                tracing::warn!(client_ip = %client_ip, "Rate limit exceeded");
                Err(actix_web::error::ErrorTooManyRequests(
                    "Rate limit exceeded. Please try again later."
                ))
            }
        })
    }
}