102 lines
2.7 KiB
Rust
102 lines
2.7 KiB
Rust
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<Mutex<HashMap<String, Arc<TokenBucket>>>>,
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
struct TokenBucket {
|
|
state: Mutex<BucketState>,
|
|
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<TokenBucket> {
|
|
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<RateLimiter>,
|
|
auth_header: Option<TypedHeader<Authorization<Bearer>>>,
|
|
req: Request,
|
|
next: Next,
|
|
) -> Result<Response, StatusCode> {
|
|
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)
|
|
}
|
|
}
|