use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use axum::extract::{ConnectInfo, Request, State};
use axum::http::StatusCode;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
pub struct RateLimiter {
max_requests: u32,
window: Duration,
buckets: Mutex<HashMap<String, (Instant, u32)>>,
}
impl RateLimiter {
pub fn new(max_requests: u32, window: Duration) -> Self {
Self {
max_requests,
window,
buckets: Mutex::new(HashMap::new()),
}
}
pub fn from_env() -> Self {
let max_requests = std::env::var("YSR_AUTH_RATE_LIMIT_MAX")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(10);
let window_secs = std::env::var("YSR_AUTH_RATE_LIMIT_WINDOW_SECS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(60);
Self::new(max_requests, Duration::from_secs(window_secs))
}
fn allow(&self, key: &str) -> bool {
let mut buckets = self.buckets.lock().expect("rate limiter mutex poisoned");
let now = Instant::now();
if buckets.len() > 128 {
let window = self.window;
buckets.retain(|_, (start, _)| now.duration_since(*start) < window);
}
let entry = buckets.entry(key.to_string()).or_insert((now, 0));
if now.duration_since(entry.0) >= self.window {
*entry = (now, 0);
}
entry.1 += 1;
entry.1 <= self.max_requests
}
}
pub async fn enforce(
State(limiter): State<std::sync::Arc<RateLimiter>>,
req: Request,
next: Next,
) -> Response {
let key = req
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ConnectInfo(addr)| addr.ip().to_string())
.unwrap_or_else(|| "unknown".to_string());
if !limiter.allow(&key) {
tracing::warn!(client = %key, path = %req.uri().path(), "auth rate limit exceeded");
return StatusCode::TOO_MANY_REQUESTS.into_response();
}
next.run(req).await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn allows_requests_within_the_limit() {
let limiter = RateLimiter::new(3, Duration::from_secs(60));
assert!(limiter.allow("1.2.3.4"));
assert!(limiter.allow("1.2.3.4"));
assert!(limiter.allow("1.2.3.4"));
}
#[test]
fn rejects_requests_past_the_limit() {
let limiter = RateLimiter::new(2, Duration::from_secs(60));
assert!(limiter.allow("1.2.3.4"));
assert!(limiter.allow("1.2.3.4"));
assert!(!limiter.allow("1.2.3.4"));
}
#[test]
fn tracks_separate_keys_independently() {
let limiter = RateLimiter::new(1, Duration::from_secs(60));
assert!(limiter.allow("1.2.3.4"));
assert!(limiter.allow("5.6.7.8"));
assert!(!limiter.allow("1.2.3.4"));
}
#[test]
fn resets_after_the_window_elapses() {
let limiter = RateLimiter::new(1, Duration::from_millis(50));
assert!(limiter.allow("1.2.3.4"));
assert!(!limiter.allow("1.2.3.4"));
std::thread::sleep(Duration::from_millis(60));
assert!(limiter.allow("1.2.3.4"));
}
#[tracing_test::traced_test]
#[tokio::test]
async fn logs_a_warning_when_the_rate_limit_is_exceeded() {
use axum::Router;
use axum::body::Body;
use axum::http::Request;
use axum::routing::get;
use tower::ServiceExt;
let limiter = std::sync::Arc::new(RateLimiter::new(1, Duration::from_secs(60)));
let app = Router::new()
.route("/probe", get(|| async { StatusCode::OK }))
.layer(axum::middleware::from_fn_with_state(limiter, enforce));
app.clone()
.oneshot(
Request::builder()
.uri("/probe")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(!logs_contain("auth rate limit exceeded"));
let response = app
.oneshot(
Request::builder()
.uri("/probe")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
assert!(logs_contain("auth rate limit exceeded"));
}
}