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;
#[derive(Debug, Clone)]
pub struct RateLimitConfig {
pub requests_per_second: u32,
pub burst_size: u32,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self::new(100)
}
}
impl RateLimitConfig {
pub fn new(requests_per_second: u32) -> Self {
Self {
requests_per_second,
burst_size: requests_per_second * 2,
}
}
}
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();
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
}
}
}
pub struct RateLimit {
config: RateLimitConfig,
buckets: Arc<Mutex<Vec<TokenBucket>>>,
}
impl RateLimit {
pub fn new(config: RateLimitConfig) -> Self {
Self {
buckets: Arc::new(Mutex::new(Vec::new())),
config,
}
}
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;
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."
))
}
})
}
}