use std::{ collections::HashMap, sync::{Arc, Mutex}, time::Instant, }; use axum::{extract::Request, http::StatusCode, middleware::Next, response::Response}; use axum_extra::{ TypedHeader, headers::{Authorization, authorization::Bearer}, }; /// Per-user token-bucket rate limiter. Each bearer token gets its own bucket /// that refills to `max_per_second` tokens every second. #[derive(Clone, Debug)] pub struct RateLimiter { max_per_second: u64, buckets: Arc>>>, } #[derive(Debug)] struct TokenBucket { state: Mutex, max_tokens: u64, } #[derive(Debug)] struct BucketState { tokens: u64, last_refill: Instant, } impl RateLimiter { /// Create a new per-user rate limiter. /// /// # Panics /// /// Panics if `max_per_second` is 0. pub fn new(max_per_second: u64) -> Self { assert!( max_per_second > 0, "max_per_second must be > 0 (set rate_limit_per_user_per_second to null in config to disable)" ); Self { max_per_second, buckets: Arc::new(Mutex::new(HashMap::new())), } } fn get_or_create_bucket(&self, token: &str) -> Arc { self.buckets .lock() .expect("rate limiter lock poisoned") .entry(token.to_owned()) .or_insert_with(|| { Arc::new(TokenBucket { state: Mutex::new(BucketState { tokens: self.max_per_second, last_refill: Instant::now(), }), max_tokens: self.max_per_second, }) }) .clone() } } impl TokenBucket { fn try_acquire(&self) -> bool { let mut state = self.state.lock().expect("token bucket lock poisoned"); let now = Instant::now(); if now.duration_since(state.last_refill).as_secs() >= 1 { state.tokens = self.max_tokens; state.last_refill = now; } if state.tokens > 0 { state.tokens -= 1; true } else { false } } } pub async fn rate_limit_middleware( axum::extract::State(limiter): axum::extract::State, auth_header: Option>>, req: Request, next: Next, ) -> Result { let Some(TypedHeader(auth)) = auth_header else { return Ok(next.run(req).await); }; let bucket = limiter.get_or_create_bucket(auth.token()); if bucket.try_acquire() { Ok(next.run(req).await) } else { Err(StatusCode::TOO_MANY_REQUESTS) } }