use crate::clock::Stamp;
use std::collections::HashMap;
use std::sync::Mutex;
use std::time::Duration;
use axum::extract::Request;
use axum::http::HeaderValue;
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use crate::Error;
use crate::auth::CurrentUser;
const SWEEP_AT: usize = 10_000;
pub(crate) struct Limiter {
id: String,
max: u32,
window: Duration,
hits: Mutex<HashMap<String, (u32, Stamp)>>,
}
pub(crate) enum Verdict {
Allowed { remaining: u32 },
Limited { retry_after: u64 },
}
impl Limiter {
pub fn new(id: String, max: u32, window: Duration) -> Self {
Self {
id,
max: max.max(1),
window,
hits: Mutex::new(HashMap::new()),
}
}
pub fn hit(&self, key: &str) -> Verdict {
let mut hits = self.hits.lock().unwrap_or_else(|e| e.into_inner());
if hits.len() >= SWEEP_AT {
hits.retain(|_, (_, start)| start.elapsed() < self.window);
}
let entry = hits.entry(key.to_owned()).or_insert((0, Stamp::now()));
if entry.1.elapsed() >= self.window {
*entry = (0, Stamp::now());
}
if entry.0 >= self.max {
let left = self.window.saturating_sub(entry.1.elapsed());
return Verdict::Limited {
retry_after: left.as_secs().max(1),
};
}
entry.0 += 1;
Verdict::Allowed {
remaining: self.max - entry.0,
}
}
pub fn max(&self) -> u32 {
self.max
}
}
fn key(req: &Request) -> String {
let user = req
.extensions()
.get::<CurrentUser>()
.and_then(|c| c.user.as_ref().map(|u| u.id));
match user {
Some(id) => format!("user:{id}"),
None => match crate::ClientIp::of(req) {
Some(ip) => format!("ip:{ip}"),
None => "ip:unknown".to_owned(),
},
}
}
async fn shared_hit(limiter: &Limiter, db: &crate::db::Db, key: &str) -> Verdict {
let key = format!("throttle:{}:{key}", limiter.id);
match crate::counters::increment(db, &key, limiter.window).await {
Ok((count, ends)) if count > limiter.max => Verdict::Limited {
retry_after: crate::counters::seconds_until(ends),
},
Ok((count, _)) => Verdict::Allowed {
remaining: limiter.max - count,
},
Err(err) => {
tracing::warn!(error = ?err, "could not count a rate-limited request");
Verdict::Allowed {
remaining: limiter.max,
}
}
}
}
pub(crate) async fn check(limiter: &Limiter, req: Request, next: Next) -> Response {
let shared = req
.extensions()
.get::<crate::AppState>()
.filter(|state| state.config.cache_store == crate::CacheStore::Database)
.map(|state| state.db.clone());
let verdict = match &shared {
Some(db) => shared_hit(limiter, db, &key(&req)).await,
None => limiter.hit(&key(&req)),
};
respond(verdict, limiter.max(), req, next).await
}
#[non_exhaustive]
pub struct LimitRequest<'a> {
pub user: Option<&'a crate::auth::User>,
pub ip: Option<std::net::IpAddr>,
pub method: &'a axum::http::Method,
pub path: &'a str,
pub headers: &'a axum::http::HeaderMap,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Limit {
max: u32,
per: Duration,
key: Option<String>,
unlimited: bool,
}
impl Limit {
pub fn per_minute(max: u32) -> Self {
Self::per(max, Duration::from_secs(60))
}
pub fn per_hour(max: u32) -> Self {
Self::per(max, Duration::from_secs(3600))
}
pub fn per(max: u32, per: Duration) -> Self {
Self {
max: max.max(1),
per: per.max(Duration::from_secs(1)),
key: None,
unlimited: false,
}
}
pub fn none() -> Self {
Self {
max: u32::MAX,
per: Duration::from_secs(1),
key: None,
unlimited: true,
}
}
pub fn by(mut self, key: impl Into<String>) -> Self {
self.key = Some(key.into());
self
}
}
pub(crate) type LimitRule = std::sync::Arc<dyn Fn(&LimitRequest) -> Limit + Send + Sync>;
pub(crate) struct NamedLimiter {
pub rule: LimitRule,
hits: Mutex<HashMap<String, (u32, Stamp)>>,
}
impl NamedLimiter {
pub fn new(rule: LimitRule) -> Self {
Self {
rule,
hits: Mutex::new(HashMap::new()),
}
}
fn hit(&self, key: &str, limit: &Limit) -> Verdict {
let mut hits = self.hits.lock().unwrap_or_else(|e| e.into_inner());
if hits.len() >= SWEEP_AT {
hits.retain(|_, (_, start)| start.elapsed() < Duration::from_secs(24 * 60 * 60));
}
let entry = hits.entry(key.to_owned()).or_insert((0, Stamp::now()));
if entry.1.elapsed() >= limit.per {
*entry = (0, Stamp::now());
}
if entry.0 >= limit.max {
let left = limit.per.saturating_sub(entry.1.elapsed());
return Verdict::Limited {
retry_after: left.as_secs().max(1),
};
}
entry.0 += 1;
Verdict::Allowed {
remaining: limit.max - entry.0,
}
}
}
pub(crate) async fn check_named(name: &str, req: Request, next: Next) -> Response {
let Some(state) = req.extensions().get::<crate::AppState>().cloned() else {
return next.run(req).await;
};
let Some(limiter) = state.limiters.get(name) else {
return Error::Internal(anyhow::anyhow!(
"no rate limiter `{name}`: define it with App::rate_limiter"
))
.into_response();
};
let user = req
.extensions()
.get::<CurrentUser>()
.and_then(|c| c.user.clone());
let limit = (limiter.rule)(&LimitRequest {
user: user.as_deref(),
ip: crate::ClientIp::of(&req),
method: req.method(),
path: req.uri().path(),
headers: req.headers(),
});
if limit.unlimited {
return next.run(req).await;
}
let key = format!(
"{name}:{}:{}",
limit.key.clone().unwrap_or_else(|| key(&req)),
limit.per.as_secs()
);
let verdict = if state.config.cache_store == crate::CacheStore::Database {
let shared = Limiter::new(format!("named:{name}"), limit.max, limit.per);
shared_hit(&shared, &state.db, &key).await
} else {
limiter.hit(&key, &limit)
};
respond(verdict, limit.max, req, next).await
}
async fn respond(verdict: Verdict, max: u32, req: Request, next: Next) -> Response {
match verdict {
Verdict::Limited { retry_after } => {
let mut res = Error::TooManyRequests.into_response();
if let Ok(value) = HeaderValue::from_str(&retry_after.to_string()) {
res.headers_mut().insert("retry-after", value);
}
res
}
Verdict::Allowed { remaining } => {
let mut res = next.run(req).await;
let headers = res.headers_mut();
headers.insert("x-ratelimit-limit", max.into());
headers.insert("x-ratelimit-remaining", remaining.into());
res
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn counts_per_key_and_window() {
let limiter = Limiter::new("t".into(), 2, Duration::from_millis(50));
assert!(matches!(
limiter.hit("a"),
Verdict::Allowed { remaining: 1 }
));
assert!(matches!(
limiter.hit("a"),
Verdict::Allowed { remaining: 0 }
));
assert!(matches!(
limiter.hit("a"),
Verdict::Limited { retry_after: 1 }
));
assert!(matches!(limiter.hit("b"), Verdict::Allowed { .. }));
std::thread::sleep(Duration::from_millis(60));
assert!(matches!(
limiter.hit("a"),
Verdict::Allowed { remaining: 1 }
));
}
#[test]
fn old_counters_are_swept_once_there_are_many() {
let limiter = Limiter::new("t".into(), 5, Duration::from_secs(60));
let named = NamedLimiter::new(std::sync::Arc::new(|_: &LimitRequest| Limit::per_hour(5)));
let limit = Limit::per_hour(5);
for i in 0..SWEEP_AT {
limiter.hit(&format!("ip:{i}"));
named.hit(&format!("ip:{i}"), &limit);
}
crate::clock::with_offset_sync(25 * 60 * 60, || {
limiter.hit("late");
named.hit("late", &limit);
});
assert_eq!(limiter.hits.lock().unwrap().len(), 1);
assert_eq!(named.hits.lock().unwrap().len(), 1);
crate::clock::with_offset_sync(-10, || {
assert!(matches!(
limiter.hit("late"),
Verdict::Allowed { remaining: 3 }
));
});
}
#[tokio::test]
async fn a_named_limit_outside_the_app_lets_requests_through_and_an_unknown_one_is_a_500() {
use tower::ServiceExt;
let router = |name: &'static str| {
axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(axum::middleware::from_fn(move |req, next| {
check_named(name, req, next)
}))
};
let request = || {
Request::builder()
.uri("/")
.body(axum::body::Body::empty())
.unwrap()
};
let res = router("api").oneshot(request()).await.unwrap();
assert_eq!(res.status(), axum::http::StatusCode::OK);
let app = crate::testing::TestApp::new(crate::App::new()).await;
let mut req = request();
req.extensions_mut().insert(app.state().clone());
let res = router("nowhere").oneshot(req).await.unwrap();
assert_eq!(res.status(), axum::http::StatusCode::INTERNAL_SERVER_ERROR);
}
}